gen3-dataops-toolkit 2.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (53) hide show
  1. g3dt/__init__.py +0 -0
  2. g3dt/cli/__init__.py +5 -0
  3. g3dt/cli/_internal/__init__.py +1 -0
  4. g3dt/cli/_internal/dispatch.py +428 -0
  5. g3dt/cli/_internal/registry.py +65 -0
  6. g3dt/cli/_internal/resolve.py +22 -0
  7. g3dt/cli/_internal/runner.py +76 -0
  8. g3dt/cli/_internal/safety.py +110 -0
  9. g3dt/cli/config_cmds.py +202 -0
  10. g3dt/cli/delete_cmds.py +101 -0
  11. g3dt/cli/dict_cmds.py +102 -0
  12. g3dt/cli/ec2_cmds.py +114 -0
  13. g3dt/cli/indexd_cmds.py +57 -0
  14. g3dt/cli/jobs.py +83 -0
  15. g3dt/cli/k8s.py +54 -0
  16. g3dt/cli/main.py +110 -0
  17. g3dt/cli/metadata.py +76 -0
  18. g3dt/cli/synth.py +206 -0
  19. g3dt/config.py +393 -0
  20. g3dt/indexd/__init__.py +0 -0
  21. g3dt/indexd/indexd_registrar.py +244 -0
  22. g3dt/ingest/ingest.py +629 -0
  23. g3dt/resolver.py +163 -0
  24. g3dt/services/delete/delete_all_metadata_for_project.py +170 -0
  25. g3dt/services/delete/delete_metadata.sh +153 -0
  26. g3dt/services/delete/delete_metadata_by_guid.py +338 -0
  27. g3dt/services/dictionary/deploy_dd.sh +65 -0
  28. g3dt/services/dictionary/pull_dict.sh +59 -0
  29. g3dt/services/dictionary/upload_dictionary.py +109 -0
  30. g3dt/services/indexd/register_indexd.py +240 -0
  31. g3dt/services/k8s_ops/argocd_restart_etl.sh +140 -0
  32. g3dt/services/k8s_ops/argocd_restart_ms.sh +102 -0
  33. g3dt/services/k8s_ops/argocd_restart_schema.sh +106 -0
  34. g3dt/services/k8s_ops/login_to_pod.sh +110 -0
  35. g3dt/services/k8s_ops/restart_etl_and_ms.sh +56 -0
  36. g3dt/services/synthetic_data/delete_synth_metadata_sheepdog.py +183 -0
  37. g3dt/services/synthetic_data/full_deploy_dd_and_synth.sh +124 -0
  38. g3dt/services/synthetic_data/generate_synth_metadata.sh +133 -0
  39. g3dt/services/synthetic_data/upload_synth_metadata_sheepdog.py +165 -0
  40. g3dt/services/upload/metadata/upload_all_studies.sh +108 -0
  41. g3dt/services/upload/metadata/upload_metadata.py +152 -0
  42. g3dt/upload/__init__.py +1 -0
  43. g3dt/upload/metadata_deleter.py +265 -0
  44. g3dt/upload/metadata_submitter.py +1093 -0
  45. g3dt/upload/upload_synthdata_s3.py +164 -0
  46. g3dt/utils/athena_utils.py +834 -0
  47. g3dt/utils/dbt_utils.py +66 -0
  48. g3dt/utils/release_writer.py +188 -0
  49. g3dt/validate/validate.py +609 -0
  50. gen3_dataops_toolkit-2.0.0.dist-info/METADATA +125 -0
  51. gen3_dataops_toolkit-2.0.0.dist-info/RECORD +53 -0
  52. gen3_dataops_toolkit-2.0.0.dist-info/WHEEL +4 -0
  53. gen3_dataops_toolkit-2.0.0.dist-info/entry_points.txt +3 -0
@@ -0,0 +1,834 @@
1
+ import logging
2
+ import boto3
3
+ import awswrangler as wr
4
+ import pandas as pd
5
+ import json
6
+ import re
7
+ import ast
8
+ from typing import Optional, Dict, Any, Union
9
+ from dataclasses import dataclass
10
+ from datetime import datetime, date
11
+ import uuid
12
+ import pytz
13
+ import base64
14
+ import numpy as np
15
+ from decimal import Decimal
16
+ from gen3_validator.dict import DataDictionary
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+ @dataclass
21
+ class AthenaConfig:
22
+ """
23
+ Configuration class for dbt release info writing.
24
+
25
+ Args:
26
+ aws_region (str): AWS region for Athena and S3 operations.
27
+ aws_profile (str): AWS profile to use for authentication.
28
+ athena_s3_output (str): S3 location for Athena query output.
29
+ """
30
+ aws_region: str
31
+ aws_profile: str
32
+ athena_s3_output: str
33
+
34
+ def as_dict(self) -> Dict[str, Any]:
35
+ return {
36
+ "aws_region": self.aws_region,
37
+ "aws_profile": self.aws_profile,
38
+ "athena_s3_output": self.athena_s3_output,
39
+ }
40
+
41
+ # ----------------- Athena Helpers -----------------
42
+
43
+ class AthenaQuery:
44
+ def __init__(self, athena_config):
45
+ self.config = athena_config
46
+
47
+ def _get_boto_session(self):
48
+ """Creates a boto3 session with the configured region and optional profile."""
49
+ region = self.config.aws_region
50
+ profile = getattr(self.config, "aws_profile", None)
51
+ logger.debug(
52
+ f"Creating boto3 session with region: {region}, profile: {profile}"
53
+ )
54
+ if not region:
55
+ logger.error("config region must be set for boto3 session.")
56
+ raise RuntimeError("config region must be set")
57
+ # Only require region; profile is optional
58
+ if profile:
59
+ session = boto3.Session(region_name=region, profile_name=profile)
60
+ else:
61
+ session = boto3.Session(region_name=region)
62
+ logger.debug("boto3 session successfully created.")
63
+ return session
64
+
65
+ def list_tables(self, database: str) -> list:
66
+ """
67
+ List all table names in the specified Athena database using awswrangler.
68
+
69
+ Args:
70
+ database (str): The Athena database to list tables from.
71
+
72
+ Returns:
73
+ list: A list of table names (str).
74
+ """
75
+ boto3_session = self._get_boto_session()
76
+ try:
77
+ tables = wr.catalog.get_tables(database=database, boto3_session=boto3_session)
78
+ table_names = [tbl['Name'] for tbl in tables]
79
+ logger.info(f"Found {len(table_names)} tables in database '{database}'.")
80
+ return table_names
81
+ except Exception as e:
82
+ logger.error(f"Error listing tables in Athena database '{database}': {e}", exc_info=True)
83
+ raise
84
+
85
+ def query_athena(self, sql: str, athena_database: str, ctas_approach: bool = True) -> pd.DataFrame:
86
+ """Runs an Athena query and returns the results as a pandas DataFrame."""
87
+ logger.info(f"Running Athena query: {sql}")
88
+ boto3_session = self._get_boto_session()
89
+ try:
90
+ df = wr.athena.read_sql_query(
91
+ sql=sql,
92
+ boto3_session=boto3_session,
93
+ database=athena_database,
94
+ ctas_approach=ctas_approach,
95
+ s3_output=self.config.athena_s3_output
96
+ )
97
+ logger.info(
98
+ f"Athena query completed successfully. Returned {len(df)} rows."
99
+ )
100
+ except Exception as e:
101
+ logger.error(f"Error running Athena query: {e}", exc_info=True)
102
+ raise
103
+ return df
104
+
105
+ def create_release_table(
106
+ self, release_db: str, release_table: str, release_s3_location: str
107
+ ) -> None:
108
+ """
109
+ Create the release tracking table in the given database if it does not already exist.
110
+
111
+ The table is created as an Iceberg table with Parquet format and Snappy compression.
112
+ The schema includes release_tag, model_name, db_name, snapshot_id, committed_at, inserted_at, and github_sha.
113
+
114
+ The names are parameters (resolved by the caller, e.g. from SSM
115
+ release/db + release/table + buckets/metadata) so this module stays a
116
+ pure, name-free helper library.
117
+
118
+ Raises:
119
+ Exception: If table creation fails.
120
+ """
121
+ create_sql = f"""
122
+ CREATE TABLE IF NOT EXISTS {release_db}.{release_table} (
123
+ release_tag STRING,
124
+ model_name STRING,
125
+ db_name STRING,
126
+ snapshot_id BIGINT,
127
+ committed_at TIMESTAMP,
128
+ inserted_at TIMESTAMP,
129
+ github_sha STRING
130
+ )
131
+ PARTITIONED BY (release_tag)
132
+ LOCATION '{release_s3_location}'
133
+ TBLPROPERTIES (
134
+ 'table_type'='ICEBERG',
135
+ 'format'='parquet',
136
+ 'write_compression'='snappy'
137
+ )
138
+ """
139
+ logger.info(f"Ensuring release table exists: {release_db}.{release_table}")
140
+ try:
141
+ self.query_athena(create_sql, release_db, False)
142
+ logger.info(f"Release table {release_db}.{release_table} created or already exists.")
143
+ except Exception as e:
144
+ logger.error(f"Failed to create release table {release_db}.{release_table}: {e}", exc_info=True)
145
+ raise
146
+
147
+ def insert_to_iceberg_table(self, df, table_name, athena_database):
148
+ """Inserts row into table."""
149
+ logger.info(
150
+ f"Inserting DataFrame with {len(df)} rows into iceberg table '{table_name}' "
151
+ f"in database '{athena_database}'."
152
+ )
153
+ boto3_session = self._get_boto_session()
154
+ try:
155
+ logger.info(f"Inserting to icerberg table with s3 output: {self.config.athena_s3_output}")
156
+ temp_s3_path = f"{self.config.athena_s3_output}/temp/{uuid.uuid4()}/"
157
+ wr.athena.to_iceberg(
158
+ df=df,
159
+ database=athena_database,
160
+ table=table_name,
161
+ boto3_session=boto3_session,
162
+ workgroup='primary',
163
+ temp_path=temp_s3_path
164
+ )
165
+ logger.info(
166
+ f"Insert to iceberg table '{table_name}' successful."
167
+ )
168
+ except Exception as e:
169
+ logger.error(
170
+ f"Error inserting to iceberg table '{table_name}': {e}",
171
+ exc_info=True
172
+ )
173
+ raise
174
+
175
+ def find_db_for_model(self, model_name: str) -> Optional[str]:
176
+ """
177
+ Search all Athena databases for a table with the given name and return the database name if found.
178
+
179
+ This function iterates through all available Athena databases and checks if a table with the specified
180
+ model_name exists in any of them. If found, it returns the name of the database; otherwise, it returns None.
181
+
182
+ Args:
183
+ model_name (str): The name of the table/model to search for.
184
+ Returns:
185
+ Optional[str]: The name of the database containing the table, or None if not found.
186
+
187
+ Raises:
188
+ Exception: If there is an error fetching databases or tables from Athena.
189
+
190
+ Example:
191
+ >>> find_db_for_model("my_table", config)
192
+ 'my_database'
193
+ >>> find_db_for_model("nonexistent_table", config)
194
+ None
195
+
196
+ Notes:
197
+ - Requires AWS credentials and permissions to list Athena databases and tables.
198
+ - Uses the awswrangler (wr) library.
199
+ - The AWS region is determined from config.aws_region.
200
+
201
+ """
202
+ try:
203
+ # Create a session and pass it to wrangler
204
+ boto3_session = self._get_boto_session()
205
+ databases = wr.catalog.databases(boto3_session=boto3_session)
206
+ db_list = databases.get('Database', [])
207
+ for db in db_list:
208
+ try:
209
+ # Pass the session to other wrangler calls too
210
+ tables = wr.catalog.get_tables(database=db, boto3_session=boto3_session)
211
+ table_names = [t.get('Name') for t in tables]
212
+ if model_name in table_names:
213
+ return db
214
+ except Exception as table_exc:
215
+ logger.warning(f"Could not fetch tables for database '{db}': {table_exc}")
216
+ continue
217
+ return None
218
+ except Exception as e:
219
+ logger.error(f"Error searching for model '{model_name}' in Athena databases: {e}")
220
+ return None
221
+
222
+
223
+ def write_iceberg_to_db(
224
+ df: pd.DataFrame,
225
+ database: str,
226
+ table: str,
227
+ athena_s3_output: str,
228
+ workgroup: str = "primary",
229
+ table_location: str = None,
230
+ partition_cols: list = None,
231
+ merge_cols: list = None,
232
+ schema_evolution: bool = False,
233
+ boto3_session=None,
234
+ ) -> None:
235
+ """
236
+ Write a DataFrame to an Iceberg table via Athena.
237
+
238
+ Parameters
239
+ ----------
240
+ df : pd.DataFrame
241
+ DataFrame to write.
242
+ database : str
243
+ Glue database name.
244
+ table : str
245
+ Glue/Athena Iceberg table name.
246
+ athena_s3_output : str
247
+ S3 URI for Athena query results and temporary staging
248
+ (e.g., 's3://bucket/athena-output/').
249
+ workgroup : str, optional
250
+ Athena workgroup to use. Defaults to "primary".
251
+ table_location : str, optional
252
+ S3 path where the Iceberg table data files are stored.
253
+ Required when creating a new table. If None, the table
254
+ must already exist.
255
+ partition_cols : list, optional
256
+ Columns to partition the Iceberg table by. If None, no
257
+ partitioning is applied.
258
+ merge_cols : list, optional
259
+ Columns to use for MERGE INTO (upsert) semantics. If None,
260
+ rows are appended.
261
+ schema_evolution : bool, optional
262
+ If True, allow schema evolution for new columns. Defaults
263
+ to False.
264
+ boto3_session : boto3.Session, optional
265
+ A boto3 session. If None, the default session is used.
266
+ """
267
+ if table_location is None:
268
+ logger.warning(
269
+ "table_location is None for %s.%s. "
270
+ "This will fail if the Iceberg table does not "
271
+ "already exist. Set 's3_path' in your config "
272
+ "to provide a table location for new tables.",
273
+ database, table,
274
+ )
275
+
276
+ try:
277
+ logger.debug(f"Creating Glue database '{database}' if not exists.")
278
+ wr.catalog.create_database(name=database, exist_ok=True)
279
+
280
+ temp_path = f"{athena_s3_output.rstrip('/')}/temp/{uuid.uuid4()}/"
281
+ logger.debug(
282
+ f"Writing DataFrame to Iceberg table {database}.{table} "
283
+ f"(workgroup: {workgroup}, temp_path: {temp_path}, "
284
+ f"table_location: {table_location})"
285
+ )
286
+ wr.athena.to_iceberg(
287
+ df=df.astype("string"),
288
+ database=database,
289
+ table=table,
290
+ temp_path=temp_path,
291
+ table_location=table_location,
292
+ partition_cols=partition_cols,
293
+ merge_cols=merge_cols,
294
+ schema_evolution=schema_evolution,
295
+ workgroup=workgroup,
296
+ boto3_session=boto3_session,
297
+ )
298
+ logger.debug(
299
+ f"Successfully wrote to Iceberg table {database}.{table}"
300
+ )
301
+ except Exception as e:
302
+ if "Must specify table location" in str(e):
303
+ logger.error(
304
+ "Iceberg table %s.%s does not exist and no "
305
+ "table_location was provided. Set 's3_path' "
306
+ "in common.metadata_upload in "
307
+ "the caller to specify where the "
308
+ "table data should be stored.",
309
+ database, table,
310
+ )
311
+ logger.error(
312
+ f"Failed to write to Iceberg table "
313
+ f"{database}.{table}: {e}"
314
+ )
315
+ raise RuntimeError(
316
+ f"Failed to write to Iceberg table "
317
+ f"{database}.{table}: {e}"
318
+ )
319
+
320
+
321
+ def convert_dataframe_types_for_json(df: pd.DataFrame) -> pd.DataFrame:
322
+ """
323
+ Converts DataFrame column types to JSON-serialisable formats.
324
+ """
325
+ df_copy = df.copy()
326
+ for col in df_copy.columns:
327
+ # Handle Decimal objects first, converting them to float
328
+ if any(isinstance(x, Decimal) for x in df_copy[col].dropna()):
329
+ df_copy[col] = df_copy[col].apply(
330
+ lambda x: float(x) if isinstance(x, Decimal) else x
331
+ )
332
+
333
+ # Handle datetime-like types, converting to ISO strings
334
+ if pd.api.types.is_datetime64_any_dtype(df_copy[col].dtype):
335
+ df_copy[col] = df_copy[col].apply(
336
+ lambda x: x.isoformat() if pd.notna(x) else None
337
+ )
338
+ continue
339
+
340
+ # For float columns (either original or from Decimal conversion)
341
+ if pd.api.types.is_float_dtype(df_copy[col].dtype):
342
+ # Only convert to object type if NaN is present
343
+ if df_copy[col].isna().any():
344
+ df_copy[col] = df_copy[col].astype(object).where(df_copy[col].notna(), None)
345
+ else:
346
+ # Ensures columns of Decimals without Nones become float
347
+ df_copy[col] = df_copy[col].astype(float)
348
+
349
+ # Convert numpy integers to standard Python integers
350
+ elif pd.api.types.is_integer_dtype(df_copy[col].dtype):
351
+ df_copy[col] = df_copy[col].astype('Int64')
352
+
353
+ # Handle object columns that might contain Timestamps or other special types
354
+ elif pd.api.types.is_object_dtype(df_copy[col].dtype):
355
+ df_copy[col] = df_copy[col].apply(
356
+ lambda x: x.isoformat() if isinstance(x, (pd.Timestamp, datetime))
357
+ else None if pd.isna(x)
358
+ else x
359
+ )
360
+
361
+ return df_copy
362
+
363
+ def replace_nan_with_none(obj):
364
+ """
365
+ Recursively traverses a dictionary or list and replaces float NaN values with None.
366
+ """
367
+ if isinstance(obj, dict):
368
+ return {k: replace_nan_with_none(v) for k, v in obj.items()}
369
+ elif isinstance(obj, list):
370
+ return [replace_nan_with_none(i) for i in obj]
371
+ # Check if it's a float and is NaN
372
+ elif isinstance(obj, float) and np.isnan(obj):
373
+ return None
374
+ else:
375
+ return obj
376
+
377
+
378
+ def json_serialiser(obj):
379
+ """
380
+ Custom JSON serialiser for handling types that aren't natively JSON-serialisable.
381
+
382
+ This improved version handles:
383
+ - Null/NA values (None, np.nan, pd.NaT), which are converted to None (JSON null).
384
+ - Date and time objects (datetime.date, datetime.datetime, pd.Timestamp),
385
+ which are converted to ISO 8601 strings.
386
+ - Decimal objects, which are converted to floats.
387
+ - NumPy numeric types (integer and floating), which are converted to
388
+ standard Python int and float types.
389
+ - Bytes objects, which are Base64-encoded into a string for safe JSON transport.
390
+ - Set and frozenset objects, which are converted to lists.
391
+ """
392
+ # This check is critical and must be first. It correctly handles
393
+ # None, np.nan, and pd.NaT, converting them to None for JSON null.
394
+ if pd.isna(obj):
395
+ return None
396
+
397
+ # Handle specific, non-native JSON types
398
+ if isinstance(obj, (datetime, pd.Timestamp, date)):
399
+ return obj.isoformat()
400
+ if isinstance(obj, Decimal):
401
+ return float(obj)
402
+ if isinstance(obj, np.integer):
403
+ return int(obj)
404
+ if isinstance(obj, np.floating):
405
+ # This handles numpy float types, including np.inf.
406
+ # np.nan is already caught by the pd.isna() check above.
407
+ return float(obj)
408
+ if isinstance(obj, bytes):
409
+ # JSON does not support bytes. A common practice is to encode them
410
+ # as a Base64 string, which is safe for JSON.
411
+ return base64.b64encode(obj).decode('utf-8')
412
+ if isinstance(obj, (set, frozenset)):
413
+ # Sets are not JSON serialisable; convert them to a list.
414
+ return list(obj)
415
+
416
+ # For any unhandled types, this will raise an error. This is the correct
417
+ # protocol for a default function used with `json.dumps()`.
418
+ raise TypeError(f"Type {type(obj)} not JSON serialisable")
419
+
420
+
421
+
422
+
423
+ class AthenaValidationWriter:
424
+ """
425
+ AthenaWriter encapsulates common operations for extracting tabular data from an Amazon Athena table,
426
+ transforming it as needed, and preparing a JSON serialization suitable for downstream processing or archiving to S3.
427
+
428
+ Responsibilities:
429
+ - Retrieve the most recent study_id and snapshot_id from the specified Athena table.
430
+ - Extract the entire contents of a table.
431
+ - Clean and format specific columns, such as decoding any stringified dictionary values for submitter_id links.
432
+ - Construct clean, consistent JSON output.
433
+
434
+ Args:
435
+ athena_config (AthenaConfig): Configuration for Athena session, including region, workgroup, and output S3 location.
436
+ db_name (str): Name of the Athena database.
437
+ table_name (str): Name of the Athena table.
438
+ """
439
+ def __init__(self, athena_config, db_name, table_name):
440
+ """
441
+ Initializes an AthenaWriter instance.
442
+
443
+ Args:
444
+ athena_config (AthenaConfig): AthenaConfig instance with AWS and Athena parameters.
445
+ db_name (str): Database name in Athena.
446
+ table_name (str): Table name in Athena.
447
+ """
448
+ self.athena_config = athena_config
449
+ self.db_name = db_name
450
+ self.table_name = table_name
451
+ self.study_id = None
452
+ self.snapshot_id = None
453
+ logger.info(f"AthenaWriter initialized with db: {db_name}, table: {table_name}")
454
+
455
+
456
+ def _get_latest_snapshot_id(self, return_commit_datetime: bool = False) -> str:
457
+ """
458
+ Retrieves the latest snapshot_id from the Iceberg/Athena table's $snapshots metadata.
459
+
460
+ Returns:
461
+ str: The most recent snapshot ID based on the newest committed_at timestamp.
462
+
463
+ Notes:
464
+ The method stores the result as self.snapshot_id for convenient reuse.
465
+ If no results are present, this will raise an IndexError.
466
+ """
467
+ # Query using CAST to avoid bringing in unsupported timestamp(3) with time zone types,
468
+ # instead cast everything to string (safe for Athena <-> pandas & avoids Hive errors)
469
+ query = f"""
470
+ SELECT
471
+ CAST(snapshot_id AS VARCHAR) AS snapshot_id,
472
+ CAST(committed_at AS VARCHAR) AS committed_at
473
+ FROM "{self.db_name}"."{self.table_name}$snapshots"
474
+ ORDER BY committed_at DESC
475
+ LIMIT 1
476
+ """
477
+ athena_query = AthenaQuery(self.athena_config)
478
+ result = athena_query.query_athena(sql=query, athena_database=self.db_name, ctas_approach=False)
479
+
480
+ if result.empty:
481
+ logger.warning(
482
+ f"No snapshot rows found for {self.db_name}.{self.table_name}"
483
+ )
484
+ self.snapshot_id = None
485
+ if return_commit_datetime:
486
+ return None, None
487
+ return None
488
+
489
+ snapshot_id = result['snapshot_id'].iloc[0]
490
+ self.snapshot_id = snapshot_id
491
+ logger.info(f"Retrieved snapshot_id: {snapshot_id}")
492
+
493
+ if return_commit_datetime:
494
+ commit_datetime = result['committed_at'].iloc[0]
495
+ return snapshot_id, commit_datetime if commit_datetime is not None else None
496
+ return snapshot_id
497
+
498
+ def _get_full_table(self):
499
+ """
500
+ Retrieves the complete contents of the Athena table, optionally at a specific snapshot.
501
+
502
+ If self.snapshot_id is set, queries the table at that snapshot id.
503
+ Otherwise, retrieves the latest/current version.
504
+
505
+ Returns:
506
+ list[dict]: Each row as a dictionary where keys are column names.
507
+
508
+ Notes:
509
+ This fetches all rows, which may have memory implications for large tables.
510
+ """
511
+ table_ref = f'"{self.db_name}"."{self.table_name}"'
512
+ if self.snapshot_id is not None:
513
+ table_ref = f'{table_ref} FOR VERSION AS OF {self.snapshot_id}'
514
+ logger.info(f"Querying table at snapshot_id: {self.snapshot_id}")
515
+ else:
516
+ logger.info("Querying table at latest/current version (no snapshot_id).")
517
+
518
+ query = f"""
519
+ SELECT *
520
+ FROM {table_ref}
521
+ """
522
+ athena_query = AthenaQuery(self.athena_config)
523
+ result = athena_query.query_athena(sql=query, athena_database=self.db_name, ctas_approach=False)
524
+ logger.info(f"Retrieved {len(result)} rows from the table.")
525
+ return result
526
+
527
+ def _format_submitter_id_value(self, submitter_id_value: str) -> Union[dict, list, str]:
528
+ """
529
+ Attempts to parse a string as a dictionary or list of dictionaries,
530
+ handling both JSON format and Python dict string representations.
531
+
532
+ This handles cases where Athena returns:
533
+ - Single dict as JSON: '{"submitter_id": "value"}'
534
+ - Single dict as Python: "{'submitter_id': 'value'}"
535
+ - List of dicts as JSON: '[{"submitter_id": "val1"}, {"submitter_id": "val2"}]'
536
+
537
+ Args:
538
+ submitter_id_value: The value from the raw Athena row.
539
+
540
+ Returns:
541
+ dict, list, or original value: Parsed dict/list if successfully parsed,
542
+ else the input value unchanged.
543
+ """
544
+ if not isinstance(submitter_id_value, str):
545
+ return submitter_id_value
546
+
547
+ # Try JSON parsing first (handles both dicts and lists with escaped quotes)
548
+ try:
549
+ parsed = json.loads(submitter_id_value)
550
+ # Validate that result is dict or list of dicts
551
+ if isinstance(parsed, dict):
552
+ return parsed
553
+ elif isinstance(parsed, list):
554
+ return parsed
555
+ except (json.JSONDecodeError, ValueError):
556
+ pass
557
+
558
+ # Fallback to ast.literal_eval for Python-style dicts/lists
559
+ dict_or_list_pattern = re.compile(r'^\s*[\[{]\s*["\']?submitter_id["\']?')
560
+ if dict_or_list_pattern.match(submitter_id_value):
561
+ try:
562
+ parsed = ast.literal_eval(submitter_id_value)
563
+ if isinstance(parsed, (dict, list)):
564
+ return parsed
565
+ except (ValueError, SyntaxError):
566
+ logger.warning(f"Could not parse string-like dictionary/list: {submitter_id_value}")
567
+ pass
568
+
569
+ return submitter_id_value
570
+
571
+
572
+
573
+ def construct_json(self):
574
+ """
575
+ Constructs a JSON string from the Athena table data, adding and normalizing
576
+ certain metadata fields (including consistent snapshot_id and study_id), and
577
+ parsing stringified dict values for proper downstream JSON compatibility.
578
+
579
+ Returns:
580
+ str: Formatted JSON string with all records, ready for saving to file or S3.
581
+
582
+ Steps:
583
+ 1. Fetch latest study_id and snapshot_id for the table.
584
+ 2. Retrieve all table rows.
585
+ 3. For each row, update/add 'snapshot_id' and 'study_id', and parse any
586
+ stringified dict fields (e.g., submitter_id link columns).
587
+ 4. Serialize rows as a pretty JSON string.
588
+ """
589
+ logger.info("Getting latest snapshot_id")
590
+ self._get_latest_snapshot_id()
591
+
592
+ logger.info("Constructing JSON data from Athena table...")
593
+ full_table = self._get_full_table()
594
+
595
+ if hasattr(full_table, "to_dict"):
596
+ full_table_json = full_table.to_dict(orient="records")
597
+ else:
598
+ logger.error("Expected a pandas DataFrame from _get_full_table(), but got type: %s", type(full_table))
599
+ raise ValueError("Unable to convert table to JSON, not a DataFrame.")
600
+
601
+ # Convert data types to JSON-serialisable formats
602
+ logger.info("Converting DataFrame types for JSON serialisation...")
603
+ full_table = convert_dataframe_types_for_json(full_table)
604
+
605
+ full_table_json = full_table.to_dict(orient="records")
606
+
607
+ for obj in full_table_json:
608
+ try:
609
+ # Convert string submitter_id link dicts to valid dict
610
+ for key, value in list(obj.items()):
611
+ obj[key] = self._format_submitter_id_value(value)
612
+ except Exception as e:
613
+ logger.warning(f"Error processing row for JSON output: {e}")
614
+
615
+ try:
616
+ full_table_json = replace_nan_with_none(full_table_json)
617
+ json_data = json.dumps(full_table_json, indent=4, default=json_serialiser)
618
+ logger.info("JSON data construction complete.")
619
+ except (TypeError, ValueError) as e:
620
+ logger.error(f"Failed to serialise JSON: {e}")
621
+ raise
622
+
623
+ return json_data
624
+
625
+
626
+ class AthenaGoldWriter(AthenaValidationWriter):
627
+ def __init__(self, athena_config, db_name, table_name):
628
+ super().__init__(athena_config, db_name, table_name)
629
+ self.study_id = None
630
+ self.snapshot_id = None
631
+ self.json_data = None
632
+ logger.info(f"AthenaGoldWriter initialized with db: {db_name}, table: {table_name}")
633
+
634
+ def construct_json(self) -> str:
635
+ """
636
+ This 'construct_json' is specialized for the AthenaGoldWriter.
637
+
638
+ For Gold tables, the expected behavior is:
639
+ - Retrieve the full gold table (via _get_full_table())
640
+ - Attempt to parse stringified dict fields just as in the parent, but may need table-specific logic
641
+ - Serialize as a pretty JSON string
642
+ """
643
+ logger.info("Getting latest snapshot_id")
644
+ self._get_latest_snapshot_id()
645
+
646
+ logger.info("Constructing JSON data from Athena GOLD table...")
647
+ full_table = self._get_full_table()
648
+
649
+ if hasattr(full_table, "to_dict"):
650
+ full_table_json = full_table.to_dict(orient="records")
651
+ else:
652
+ logger.error("Expected a pandas DataFrame from _get_full_table(), but got type: %s", type(full_table))
653
+ raise ValueError("Unable to convert table to JSON, not a DataFrame.")
654
+
655
+ # Convert data types to JSON-serialisable formats
656
+ logger.info("Converting DataFrame types for JSON serialisation (Gold)...")
657
+ full_table = convert_dataframe_types_for_json(full_table)
658
+ full_table_json = full_table.to_dict(orient="records")
659
+
660
+ for obj in full_table_json:
661
+ try:
662
+ # Convert string submitter_id link dicts to valid dict, if present
663
+ for key, value in list(obj.items()):
664
+ obj[key] = self._format_submitter_id_value(value)
665
+ except Exception as e:
666
+ logger.warning(f"Error processing GOLD row for JSON output: {e}")
667
+
668
+ try:
669
+ full_table_json = replace_nan_with_none(full_table_json)
670
+ json_data = json.dumps(full_table_json, indent=4, default=json_serialiser)
671
+ logger.info("GOLD JSON data construction complete.")
672
+ self.json_data = json_data
673
+ except (TypeError, ValueError) as e:
674
+ logger.error(f"Failed to serialise GOLD JSON: {e}")
675
+ raise
676
+
677
+ return json_data
678
+
679
+
680
+ def generate_validation_id():
681
+ """
682
+ Generates a unique validation identifier string, based on the current date and time in Australian Eastern Time.
683
+
684
+ Returns:
685
+ str: A string formatted as "YYYYMMDDHHMMSS", e.g., "20240531173010".
686
+ """
687
+ # Get Australian Eastern timezone (automatically handles AEST/AEDT)
688
+ australian_tz = pytz.timezone('Australia/Melbourne')
689
+
690
+ # Get current time in Australian Eastern timezone
691
+ current_date_time = datetime.now(australian_tz).strftime("%Y%m%d%H%M%S")
692
+ validation_id = f"{current_date_time}"
693
+ logger.info(f"Generated validation_id (Australian Eastern Time): {validation_id}")
694
+ return validation_id
695
+
696
+ def write_validation_json_to_s3(s3_bucket,
697
+ study_id,
698
+ validation_id,
699
+ table_name,
700
+ snapshot_id,
701
+ json_data):
702
+ """
703
+ Writes the supplied JSON data to an S3 bucket using the specified logical path and metadata.
704
+
705
+ Args:
706
+ s3_bucket (str): Name of the S3 bucket.
707
+ study_id (str): Study identifier to include in the S3 key/path.
708
+ validation_id (str): Unique validation run identifier string.
709
+ table_name (str): Source Athena table name.
710
+ snapshot_id (str): Table snapshot ID (for versioning).
711
+ json_data (str): The JSON-serialized string to be uploaded.
712
+
713
+ The S3 object key is formatted as:
714
+ validation/study_id=<study_id>/validation_id=<validation_id>/table_name=<table_name>/snapshot_id=<snapshot_id>/<table_name>.json
715
+
716
+ Side Effects:
717
+ Uploads (puts) the JSON data to the given S3 bucket.
718
+
719
+ Raises:
720
+ Any exception raised by boto3.client('s3').put_object.
721
+
722
+ Logging:
723
+ Logs both intent and completion of S3 upload.
724
+ """
725
+ s3 = boto3.client('s3')
726
+ s3_object_key = f"validation/study_id={study_id}/validation_id={validation_id}/table_name={table_name}/snapshot_id={snapshot_id}/{table_name}.json"
727
+ logger.info(f"Writing JSON data to S3 bucket: {s3_bucket}, object key: {s3_object_key}")
728
+ s3.put_object(Body=json_data, Bucket=s3_bucket, Key=s3_object_key)
729
+ logger.info(f"Object created at s3://{s3_bucket}/{s3_object_key}")
730
+
731
+
732
+
733
+ def write_gold_json_to_s3(
734
+ s3_bucket,
735
+ study_id,
736
+ table_name,
737
+ snapshot_id,
738
+ json_data,
739
+ ):
740
+ """
741
+ Write JSON data for a "Gold" Athena table to S3 in a unique and logical path.
742
+
743
+ This function uploads the supplied JSON string to the given S3 bucket, placing it under a "validation"
744
+ directory structure specific to "gold" Athena tables. The S3 object key removes the "gold_" prefix from
745
+ the table name for the actual file name and omits any validation ID in the path.
746
+
747
+ Args:
748
+ s3_bucket (str): S3 bucket name to upload to.
749
+ study_id (str): The unique identifier for the study; appears in the S3 path.
750
+ table_name (str): Athena gold table name (e.g., "gold_diagnosis"); used in path and for filename.
751
+ snapshot_id (str): Identifier for the data snapshot/version used in the path.
752
+ json_data (str): JSON string (already serialized), to upload.
753
+
754
+ S3 Path Pattern:
755
+ validation/study_id=<study_id>/table_name=<table_name>/snapshot_id=<snapshot_id>/<stripped_table_name>.json
756
+
757
+ - <stripped_table_name> is the table_name with the "gold_" prefix removed.
758
+
759
+ Example:
760
+ gold_table_name = "gold_diagnosis"
761
+ table_name = gold_table_name
762
+ Actual file will be e.g.:
763
+ validation/study_id=foo/table_name=gold_diagnosis/snapshot_id=bar/diagnosis.json
764
+
765
+ Side Effects:
766
+ Uploads JSON to S3 using boto3.
767
+
768
+ Raises:
769
+ Any exception from boto3's client('s3').put_object.
770
+
771
+ Logging:
772
+ Logs before and after the S3 put operation, including path details for traceability.
773
+ """
774
+ s3 = boto3.client('s3')
775
+
776
+ filename = table_name
777
+ if filename.startswith("gold_"):
778
+ filename = filename.replace("gold_", "")
779
+
780
+ # Remove study_id from filename if present
781
+ if study_id in filename:
782
+ filename = filename.replace(f"{study_id}_", "")
783
+ else:
784
+ logger.warning(f"Filename {filename} does not contain study_id {study_id}, writing filename as {filename}")
785
+
786
+ s3_object_key = f"gold_jsons/study_id={study_id}/table_name={table_name}/snapshot_id={snapshot_id}/{filename}.json"
787
+ logger.info(f"Writing JSON data to S3 bucket: {s3_bucket}, object key: {s3_object_key}")
788
+ s3.put_object(Body=json_data, Bucket=s3_bucket, Key=s3_object_key)
789
+ logger.info(f"Object created at s3://{s3_bucket}/{s3_object_key}")
790
+
791
+
792
+ def construct_data_import_order(s3_uri) -> list:
793
+ from g3dt.validate.validate import load_schema_from_s3_uri
794
+ schema_dict = load_schema_from_s3_uri(s3_uri)
795
+ dd = DataDictionary(schema_dict)
796
+ dd.schema = schema_dict
797
+ dd.calculate_node_order()
798
+ return dd.node_order
799
+
800
+ def write_release_jsons_to_s3(s3_bucket, release_id, study_id, table_name, json_data):
801
+ """
802
+ Write a JSON string to a specific S3 location for a given release and study.
803
+
804
+ Args:
805
+ s3_bucket (str): The S3 bucket where the file will be uploaded.
806
+ release_id (str): Release identifier used in the S3 key path.
807
+ study_id (str): Study identifier used in the S3 key path.
808
+ table_name (str): Table name (used for naming the .json file).
809
+ json_data (str): JSON data (as a string) to be uploaded.
810
+
811
+ Returns:
812
+ str: The output directory path in S3 where the file was written.
813
+
814
+ Raises:
815
+ Exception: Any exception raised by boto3.client('s3').put_object.
816
+
817
+ Example:
818
+ >>> output_dir = write_release_jsons_to_s3('my-bucket', 'release123', 'study1', 'gold_foo', '{"x":1}')
819
+ """
820
+ s3 = boto3.client('s3')
821
+
822
+ filename = table_name
823
+ study_name = study_id.split("/")[-1]
824
+ if filename.startswith(f"{study_name}_"):
825
+ filename = filename[len(f"{study_name}_"):]
826
+ else:
827
+ logger.warning(f"Filename {filename} does not start with study_name {study_name}, writing filename as {filename}")
828
+
829
+ output_dir = f"release_jsons/{release_id}/{study_id}"
830
+ s3_object_key = f"{output_dir}/{filename}.json"
831
+ logger.info(f"Writing JSON data to S3 bucket: {s3_bucket}, object key: {s3_object_key}")
832
+ s3.put_object(Body=json_data, Bucket=s3_bucket, Key=s3_object_key)
833
+ logger.info(f"Object created at s3://{s3_bucket}/{s3_object_key}")
834
+ return output_dir