Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 32 additions & 37 deletions usaspending_api/download/filestreaming/download_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
from usaspending_api.download.helpers import verify_requested_columns_available
from usaspending_api.download.helpers import write_to_download_log as write_to_log
from usaspending_api.download.helpers.cleanup_helpers import cleanup_download_files, cleanup_previous_download_attempt
from usaspending_api.download.helpers.psql_helpers import build_psql_env, run_psql_to_file
from usaspending_api.download.lookups import FILE_FORMATS, JOB_STATUS_DICT, VALUE_MAPPINGS
from usaspending_api.download.models.download_job import DownloadJob
from usaspending_api.download.models.download_job_lookup import DownloadJobLookup
Expand Down Expand Up @@ -895,55 +896,50 @@ def execute_psql(temp_sql_file_path: str, source_path: str, download_job: Downlo
"""Executes a single PSQL command within its own Subprocess"""
download_sql = Path(temp_sql_file_path).read_text()
if download_sql.startswith("\\COPY"):
# Trace library parses the SQL, but cannot understand the psql-specific \COPY command. Use standard COPY here.
download_sql = download_sql[1:]

# Stack 3 context managers: (1) psql code, (2) Download replica query, (3) (same) Postgres query
subprocess_trace = SubprocessTrace(
name=f"job.{JOB_TYPE}.download.psql",
kind=SpanKind.INTERNAL,
service="bulk-download",
)

with subprocess_trace as span:
span.set_attributes(
{
"service": "bulk-download",
"resource": str(download_sql),
"span_type": "Internal",
"source_path": str(source_path),
# download job details
"download_job_id": str(download_job.download_job_id),
"download_job_status": str(download_job.job_status.name),
"download_file_name": str(download_job.file_name),
"download_file_size": download_job.file_size if download_job.file_size is not None else 0,
"number_of_rows": download_job.number_of_rows if download_job.number_of_rows is not None else 0,
"number_of_columns": (
download_job.number_of_columns if download_job.number_of_columns is not None else 0
),
"error_message": download_job.error_message if download_job.error_message else "",
"monthly_download": str(download_job.monthly_download),
"json_request": str(download_job.json_request) if download_job.json_request else "",
}
)
span.set_attributes({
"service": "bulk-download",
"resource": str(download_sql),
"span_type": "Internal",
"source_path": str(source_path),
"download_job_id": str(download_job.download_job_id),
"download_job_status": str(download_job.job_status.name),
"download_file_name": str(download_job.file_name),
"download_file_size": download_job.file_size if download_job.file_size is not None else 0,
"number_of_rows": download_job.number_of_rows if download_job.number_of_rows is not None else 0,
"number_of_columns": download_job.number_of_columns if download_job.number_of_columns is not None else 0,
"error_message": download_job.error_message if download_job.error_message else "",
"monthly_download": str(download_job.monthly_download),
"json_request": str(download_job.json_request) if download_job.json_request else "",
})

try:
log_time = time.perf_counter()
temp_env = os.environ.copy()
if download_job and not download_job.monthly_download:
# Since terminating the process isn't guaranteed to end the DB statement,
# add timeout to client connection
temp_env["PGOPTIONS"] = (
f"--statement-timeout={settings.DOWNLOAD_DB_TIMEOUT_IN_HOURS}h "
f"--work-mem={settings.DOWNLOAD_DB_WORK_MEM_IN_MB}MB"
)

cat_command = subprocess.Popen(["cat", temp_sql_file_path], stdout=subprocess.PIPE)
subprocess.check_output(
["psql", "-q", "-o", source_path, retrieve_db_string(), "-v", "ON_ERROR_STOP=1"],
stdin=cat_command.stdout,
stderr=subprocess.STDOUT,
env=temp_env,
# Build PostgreSQL environment using helper
psql_env = build_psql_env(
dsn=retrieve_db_string(),
statement_timeout_hours=settings.DOWNLOAD_DB_TIMEOUT_IN_HOURS if (
download_job and not download_job.monthly_download) else None,
work_mem_mb=settings.DOWNLOAD_DB_WORK_MEM_IN_MB if (
download_job and not download_job.monthly_download) else None
)

# Execute psql using helper
run_psql_to_file(
sql_path=temp_sql_file_path,
output_path=source_path,
env=psql_env,
quiet=True,
on_error_stop=True
)

duration = time.perf_counter() - log_time
Expand All @@ -956,7 +952,6 @@ def execute_psql(temp_sql_file_path: str, source_path: str, download_job: Downlo
raise e
except Exception as e:
if not settings.IS_LOCAL:
# Not logging the command as it can contain the database connection string
e.cmd = "[redacted psql command]"
write_to_log(message=e, is_error=True, download_job=download_job)
sql = subprocess.check_output(["cat", temp_sql_file_path]).decode()
Expand Down
145 changes: 145 additions & 0 deletions usaspending_api/download/helpers/psql_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,145 @@
import os
import subprocess
from typing import Optional
from urllib.parse import urlparse


def build_psql_env(
dsn: str,
statement_timeout_hours: Optional[int] = None,
work_mem_mb: Optional[int] = None,
base_env: Optional[dict] = None
) -> dict:
"""Build PostgreSQL environment variables from a database connection string."""
import logging

logger = logging.getLogger(__name__)

if not dsn:
raise ValueError("DSN cannot be empty")

logger.info(f"Parsing DSN: {dsn[:30]}...")

db_url = urlparse(dsn)

env = (base_env or os.environ).copy()

# Set PostgreSQL connection parameters
env["PGHOST"] = db_url.hostname or "localhost"
env["PGPORT"] = str(db_url.port or 5432)
env["PGUSER"] = db_url.username or "postgres"
env["PGPASSWORD"] = db_url.password or ""
env["PGDATABASE"] = db_url.path.lstrip('/') if db_url.path else "postgres"

logger.info(
f"Set PGHOST={env['PGHOST']}, PGPORT={env['PGPORT']}, PGUSER={env['PGUSER']}, PGDATABASE={env['PGDATABASE']}")

# Set optional PostgreSQL options
if statement_timeout_hours or work_mem_mb:
options = []
if statement_timeout_hours:
options.append(f"--statement-timeout={statement_timeout_hours}h")
if work_mem_mb:
options.append(f"--work-mem={work_mem_mb}MB")
env["PGOPTIONS"] = " ".join(options)

return env


def run_psql_to_file(
sql_path: str,
output_path: str,
env: dict,
quiet: bool = True,
on_error_stop: bool = True
) -> None:
"""
Execute a psql command that reads SQL from a file and writes output to another file.
"""
import logging
logger = logging.getLogger(__name__)

# Log the SQL file contents for debugging
try:
with open(sql_path, 'r') as f:
sql_content = f.read()
logger.info(f"SQL file contents (first 500 chars): {sql_content[:500]}")
except Exception as e:
logger.error(f"Could not read SQL file: {e}")

psql_args = ["psql"]

if quiet:
psql_args.append("-q")

psql_args.extend(["-o", output_path])

if on_error_stop:
psql_args.extend(["-v", "ON_ERROR_STOP=1"])

logger.info(f"psql command: {' '.join(psql_args)}")
logger.info(
f"Environment: PGHOST={env.get('PGHOST')}, "
f"PGPORT={env.get('PGPORT')}, "
f"PGUSER={env.get('PGUSER')}, "
f"PGDATABASE={env.get('PGDATABASE')}"
)

# Test database connection first
logger.info("Testing database connection...")
test_process = subprocess.run(
["psql", "-c", "SELECT 1;"],
env=env,
capture_output=True,
timeout=5
)
if test_process.returncode != 0:
logger.error(f"Database connection test failed: {test_process.stderr.decode()}")
raise Exception(f"Cannot connect to database: {test_process.stderr.decode()}")
logger.info("Database connection test successful")

logger.info("Starting cat and psql processes...")

# Start cat process
cat_process = subprocess.Popen(["cat", sql_path], stdout=subprocess.PIPE, stderr=subprocess.PIPE)

# Start psql process with cat's stdout as stdin
psql_process = subprocess.Popen(
psql_args,
stdin=cat_process.stdout,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE, # Changed to PIPE to capture stderr separately
env=env,
)

# Close cat's stdout in parent so psql gets EOF when cat exits
cat_process.stdout.close()

logger.info("Waiting for processes to complete...")

# Wait for both processes to complete with timeout
try:
psql_output, psql_error = psql_process.communicate(timeout=30) # 30 second timeout
cat_process.wait(timeout=5)
except subprocess.TimeoutExpired:
logger.error("Process timed out! Killing processes...")
psql_process.kill()
cat_process.kill()
raise Exception("psql process timed out after 30 seconds") from None

logger.info(f"psql return code: {psql_process.returncode}")
logger.info(f"psql stdout: {psql_output.decode() if psql_output else 'empty'}")
logger.info(f"psql stderr: {psql_error.decode() if psql_error else 'empty'}")

# Check for errors
if psql_process.returncode != 0:
error_msg = psql_error.decode() if psql_error else psql_output.decode() if psql_output else "Unknown error"
logger.error(f"psql failed: {error_msg}")
raise subprocess.CalledProcessError(
psql_process.returncode,
psql_args,
output=psql_output,
stderr=psql_error
)

logger.info("psql completed successfully")
Original file line number Diff line number Diff line change
@@ -1,8 +1,6 @@
import logging
import os
import re
import shutil
import subprocess
import tempfile
from datetime import date, datetime

Expand All @@ -14,16 +12,11 @@
from django.db.models import Case, CharField, F, Q, Value, When

from usaspending_api.awards.v2.lookups.lookups import all_award_types_mappings as all_ats_mappings
from usaspending_api.common.csv_helpers import count_rows_in_delimited_file
from usaspending_api.common.helpers.orm_helpers import generate_raw_quoted_query
from usaspending_api.common.helpers.s3_helpers import multipart_upload
from usaspending_api.config import CONFIG
from usaspending_api.download.filestreaming.download_generation import (
apply_annotations_to_sql,
split_and_zip_data_files,
)
from usaspending_api.download.filestreaming.download_source import DownloadSource
from usaspending_api.download.helpers import pull_modified_agencies_cgacs
from usaspending_api.download.helpers.psql_helpers import build_psql_env, run_psql_to_file
from usaspending_api.download.lookups import VALUE_MAPPINGS
from usaspending_api.references.models import SubtierAgency, ToptierAgency

Expand Down Expand Up @@ -85,7 +78,6 @@ def download(self, award_type: str, agency: str | dict = "all", generate_since:
)
source.query_paths.update({"correction_delete_ind": award_map["correction_delete_ind"]})
if award_type == "Contracts":
# Add the agency_id column to the mappings
source.query_paths.update({"agency_id": "transaction__contract_data__agency_id"})
source.query_paths.move_to_end("agency_id", last=False)
source.query_paths.move_to_end("correction_delete_ind", last=False)
Expand Down Expand Up @@ -120,12 +112,12 @@ def download(self, award_type: str, agency: str | dict = "all", generate_since:

source.queryset = source.queryset.filter(Q(update_date_filter | Q(transaction__transactiondelta__isnull=False)))

# Generate file
# Generate file using helper functions
file_path = self.create_local_file(award_type, source, agency_code, generate_since)

if file_path is None:
logger.info("No new, modified, or deleted data; discarding file")
elif not settings.IS_LOCAL:
# Upload file to S3 and delete local version
logger.info("Uploading file to S3 bucket and deleting local copy")
multipart_upload(
CONFIG.MONTHLY_DOWNLOAD_S3_BUCKET_NAME,
Expand All @@ -140,9 +132,19 @@ def download(self, award_type: str, agency: str | dict = "all", generate_since:
)

def create_local_file(
self, award_type: str, source: pd.DataFrame, agency_code: str, generate_since: str | None
self, award_type: str, source: DownloadSource, agency_code: str, generate_since: str | None
) -> str | None:
"""Generate complete file from SQL query and S3 bucket deletion files, then zip it locally"""
import shutil
import subprocess

from usaspending_api.common.csv_helpers import count_rows_in_delimited_file
from usaspending_api.common.helpers.orm_helpers import generate_raw_quoted_query
from usaspending_api.download.filestreaming.download_generation import (
apply_annotations_to_sql,
split_and_zip_data_files,
)

logger.info("Generating CSV file with creations and modifications")

# Create file paths and working directory
Expand All @@ -155,30 +157,55 @@ def create_local_file(
source_path = os.path.join(working_dir, "{}.csv".format(source_name))

# Create a unique temporary file with the raw query
raw_quoted_query = generate_raw_quoted_query(source.row_emitter(None)) # None requests all headers

raw_quoted_query = generate_raw_quoted_query(source.row_emitter(None))
csv_query_annotated = apply_annotations_to_sql(raw_quoted_query, source.human_names)

(temp_sql_file, temp_sql_file_path) = tempfile.mkstemp(prefix="bd_sql_", dir="/tmp")
with open(temp_sql_file_path, "w") as file:
file.write("\\copy ({}) To STDOUT with CSV HEADER".format(csv_query_annotated))

logger.info("Generated temp SQL file {}".format(temp_sql_file_path))
# Generate the csv with \copy
cat_command = subprocess.Popen(["cat", temp_sql_file_path], stdout=subprocess.PIPE)

try:
subprocess.check_output(
["psql", "-o", source_path, os.environ["DOWNLOAD_DATABASE_URL"], "-v", "ON_ERROR_STOP=1"],
stdin=cat_command.stdout,
stderr=subprocess.STDOUT,
# Get database URL from settings or environment variable (for tests)
db_url = os.environ.get("DOWNLOAD_DATABASE_URL") or settings.DOWNLOAD_DATABASE_URL

if not db_url:
raise ValueError("DOWNLOAD_DATABASE_URL is not configured")

logger.info(f"Using database URL: {db_url[:20]}...") # Log first 20 chars for debugging

# Build PostgreSQL environment using helper
psql_env = build_psql_env(
dsn=db_url,
statement_timeout_hours=settings.DOWNLOAD_DB_TIMEOUT_IN_HOURS,
work_mem_mb=settings.DOWNLOAD_DB_WORK_MEM_IN_MB
)

logger.info(
f"Built psql environment with PGHOST={psql_env.get('PGHOST')}, PGDATABASE={psql_env.get('PGDATABASE')}")

# Execute psql using helper
run_psql_to_file(
sql_path=temp_sql_file_path,
output_path=source_path,
env=psql_env,
quiet=False,
on_error_stop=True
)

except subprocess.CalledProcessError as e:
logger.exception(e.output)
logger.exception(e.output if hasattr(e, 'output') else str(e))
raise e
finally:
# Always cleanup temp SQL file
os.close(temp_sql_file)
os.remove(temp_sql_file_path)

# Append deleted rows to the end of the file
if not self.debugging_skip_deleted:
self.add_deletion_records(source_path, working_dir, award_type, agency_code, source, generate_since)

if count_rows_in_delimited_file(source_path, has_header=True, safe=True) > 0:
# Split the CSV into multiple files and zip it up
zipfile_path = "{}{}.zip".format(settings.CSV_LOCAL_PATH, source_name)
Expand All @@ -188,8 +215,6 @@ def create_local_file(
else:
zipfile_path = None

os.close(temp_sql_file)
os.remove(temp_sql_file_path)
shutil.rmtree(working_dir)

return zipfile_path
Expand Down
Loading
Loading