diff --git a/.github/workflows/basic_tests.yml b/.github/workflows/basic_tests.yml index 33fed6a8..5eb81549 100644 --- a/.github/workflows/basic_tests.yml +++ b/.github/workflows/basic_tests.yml @@ -11,7 +11,7 @@ jobs: runs-on: ubuntu-latest strategy: matrix: - python-version: [ "3.9", "3.10", "3.11", "3.12", "3.13", "3.14" ] + python-version: [ "3.10", "3.11", "3.12", "3.13", "3.14" ] steps: - name: Check out repository diff --git a/.gitignore b/.gitignore index ff506a58..4aa9d291 100644 --- a/.gitignore +++ b/.gitignore @@ -103,6 +103,7 @@ celerybeat.pid # Environments .env +.Renviron .venv env/ venv/ @@ -178,4 +179,12 @@ debug_app/user_overrides.json debug_app/test_results.json .test_baseline.json -.test_final.json \ No newline at end of file +.test_final.json + +# Benchmark outputs (generated by running benchmarks) +benchmark_output/ +circepy_benchmarks/ +renv.lock +eunomia_data/ +renv/ +.Rprofile \ No newline at end of file diff --git a/CLAUDE.md b/CLAUDE.md index 50a1d490..c60dada7 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -32,6 +32,35 @@ git pre-commit run --all-files If pre-commit checks fail, fix the issues and re-run until they pass. +## Ibis Execution Layer: NEVER use Python in-memory operations + +The datasets this software processes are large (often 100M+ rows). Operations that pull data into Python memory will crash the process. All data processing MUST remain as lazy ibis expressions executed on the database backend. + +### Forbidden patterns in production code (`circe/execution/` and `circe/cohort_definition_set/`): + +| Pattern | Example (NEVER do this) | Instead | +|---|---|---| +| `.execute()` | `table.execute()` loads entire table into a pandas DataFrame in memory | Compose ibis expressions; let the backend execute the full query | +| `.to_pandas()` | `table.to_pandas()` pulls result set into Python | Use ibis expressions; only call `.execute()` for small scalars (e.g., `table.limit(1).count().execute()`) | +| Python iteration over results | `for row in table.select(...).distinct().to_pandas().itertuples()` | Push aggregation/distinct into ibis; use window functions or joins | +| `ibis.memtable()` with large DataFrames | Constructing a large `pd.DataFrame` and passing to `ibis.memtable()` | Read directly from the database table (passed tables already exist in the backend) | +| Loading files into Python | `pd.read_csv(...)`, reading Parquet into memory | Use ibis to read files: `ibis.read_csv()`, `ibis.read_parquet()` | + +### Existing violations in production code (DO NOT FIX — examples for reference): + +1. **`circe/cohort_definition_set/_checksum_store.py`** — uses `pandas`, `.execute()`, `pd.DataFrame()`, row iteration — should use ibis expressions end-to-end +2. **`circe/execution/engine/custom_era.py:86`** — `.execute().iloc[:, 0]` to pull concept IDs into a Python tuple +3. **`circe/execution/engine/group_demographics.py:97`** — `.to_pandas().itertuples()` to iterate over distinct concept IDs +4. **`circe/execution/ibis/operations.py:86`** — `.execute()` to check if rows exist (use `table.limit(1).count()` instead) +5. **`benchmarks/compare_cohort_outputs.py`** — full table `.execute()`, pandas row iteration, set comparison in memory + +### Allowed uses of `.execute()`: + +- **Tests only** — tests run against small in-memory DuckDB databases with tiny fixtures. Assertions on small result sets are fine. +- **Scalar values** — getting a single count or checking existence: `table.count().execute()`, `table.limit(1).execute()` (only returns 1 row) + +When writing new production code, if you find yourself reaching for `.execute()`, `.to_pandas()`, or Python iteration over ibis results, **stop** — the query can be rewritten as a lazy ibis expression. + ## Git Workflow - Do not run `git commit` — the user will handle commits - Run pre-commit checks to validate code quality before marking tasks complete diff --git a/circe/api.py b/circe/api.py index 4c8f3a56..04e2e6ad 100644 --- a/circe/api.py +++ b/circe/api.py @@ -9,8 +9,16 @@ - cohort_print_friendly(): Generate Markdown from cohort expression """ -from typing import TYPE_CHECKING, Any, Literal, Optional - +from typing import TYPE_CHECKING, Any, Literal + +from .cohort_definition_set import ( # noqa: F401 + CohortDefinition, + CohortDefinitionSet, + CohortGenerationResult, + async_generate_cohort_set, + generate_cohort_set, + summarise_generation_results, +) from .cohortdefinition import ( BuildExpressionQueryOptions, CohortExpression, @@ -114,7 +122,7 @@ def cohort_expression_from_yaml(yaml_str: str) -> CohortExpression: def build_cohort_query( expression: CohortExpression, - options: Optional[BuildExpressionQueryOptions] = None, + options: BuildExpressionQueryOptions | None = None, ) -> str: """Generate SQL query from a cohort expression. @@ -147,8 +155,8 @@ def build_cohort( *, backend: IbisBackendLike, cdm_schema: str, - vocabulary_schema: Optional[str] = None, - results_schema: Optional[str] = None, + vocabulary_schema: str | None = None, + results_schema: str | None = None, ) -> Table: """Build a cohort as a relational table expression. @@ -199,8 +207,8 @@ def write_cohort( cdm_schema: str, cohort_table: str, cohort_id: int, - vocabulary_schema: Optional[str] = None, - results_schema: Optional[str] = None, + vocabulary_schema: str | None = None, + results_schema: str | None = None, if_exists: Literal["fail", "replace"] = "fail", ) -> None: """Build and write an OHDSI cohort table. @@ -260,14 +268,12 @@ def write_cohort( def cohort_print_friendly( expression: CohortExpression, - concept_sets: Optional[list[ConceptSet]] = None, - title: Optional[str] = None, + concept_sets: list[ConceptSet] | None = None, + title: str | None = None, include_concept_sets: bool = False, ) -> str: """Generate human-readable Markdown from a cohort expression. - This is equivalent to R CirceR's `cohortPrintFriendly()` function. - Args: expression: CohortExpression instance concept_sets: Optional list of concept sets (uses expression.concept_sets if None) diff --git a/circe/chat.py b/circe/chat.py deleted file mode 100644 index fd58d8d1..00000000 --- a/circe/chat.py +++ /dev/null @@ -1,256 +0,0 @@ -""" -Chat module for interacting with LLMs to generate cohort definitions. -""" - -import json -import os -import re -import sys -from pathlib import Path -from typing import Optional - -from circe.prompt_builder import CohortPromptBuilder, ConceptSet - - -def chat_command(args): - """ - Entry point for the chat command. - """ - start_chat( - model=args.model, - prompt_type=args.prompt_type, - output=args.output, - concept_sets_file=args.concept_sets, - input_file=args.input_file, - ) - return 0 - - -def start_chat( - model: Optional[str], - prompt_type: str, - output: Optional[str], - concept_sets_file: Optional[str], - input_file: Optional[str] = None, -): - """ - Start the interactive chat session. - """ - # Check dependencies - try: - import litellm - from dotenv import load_dotenv - except ImportError: - print( - "Error: 'litellm' and 'python-dotenv' are required for chat functionality.", - file=sys.stderr, - ) - print( - "Please install them with: pip install litellm python-dotenv", - file=sys.stderr, - ) - return 1 - - # Load environment variables - load_dotenv() - - # Determine model - if not model: - model = os.getenv("LLM_MODEL", "gpt-4o") - # Handle optional temperature if needed, but litellm handles it or we pass it - - print("🚀 Starting Circe Chat") - print(f" Model: {model}") - print(f" Prompt: {prompt_type}") - print("-" * 50) - - # Load concept sets if provided - concept_sets_data = [] - if concept_sets_file: - try: - with open(concept_sets_file) as f: - raw_data = json.load(f) - # Expecting list of dicts with id, name - for item in raw_data: - concept_sets_data.append( - ConceptSet( - id=item.get("id"), - name=item.get("name"), - description=item.get("description"), - ) - ) - print(f" Loaded {len(concept_sets_data)} concept sets from {concept_sets_file}") - except Exception as e: - print(f"Error loading concept sets: {e}", file=sys.stderr) - return 1 - - # Initialize builder - builder = CohortPromptBuilder() - - try: - system_prompt = builder.load_system_prompt(prompt_type) - except Exception as e: - print(f"Error loading system prompt: {e}", file=sys.stderr) - return 1 - - # Add inference instruction if no concept sets provided - if not concept_sets_data: - system_prompt += ( - "\n\nIMPORTANT: No concept sets were provided.\n" - "You MUST infer appropriate concept sets from the clinical description.\n" - "1. Define them using `circe.vocabulary.concept_set`.\n" - "2. Add them to the builder using `.with_concept_sets(...)`.\n" - "3. Use valid OMOP Concept IDs (or realistic placeholders if exact IDs are unknown)." - ) - - messages = [{"role": "system", "content": system_prompt}] - - print("\nPlease describe the cohort you want to build (or type 'quit' to exit):") - - first_turn = True - initial_input = None - - if input_file: - try: - initial_input = Path(input_file).read_text() - print(f" Loaded clinical description from {input_file}") - except Exception as e: - print(f"Error reading input file: {e}", file=sys.stderr) - return 1 - - while True: - try: - if first_turn and initial_input: - user_input = initial_input - print("\n> [Processing input from file...]") - else: - user_input = input("\n> ") - except (EOFError, KeyboardInterrupt): - print("\nExiting chat.") - break - - if user_input.lower() in ("quit", "exit"): - break - - if not user_input.strip(): - continue - - # Turn off first_turn flag after we have a valid input - if first_turn: - first_turn = False - - # Construct user message - if len(messages) == 1: - # First user message - format nicely - formatted_content = f"\n---\n## User Task\n**Clinical Description:**\n{user_input}\n" - if concept_sets_data: - formatted_content += builder.format_concept_sets(concept_sets_data) - else: - formatted_content += "\nNo pre-defined concept sets provided. Please infer them." - - messages.append({"role": "user", "content": formatted_content}) - else: - messages.append({"role": "user", "content": user_input}) - - # Call AI - print("Thinking...") - try: - response = litellm.completion(model=model, messages=messages) - content = response.choices[0].message.content - print("\n" + content) - - messages.append({"role": "assistant", "content": content}) - - # Extract and process code - _process_response_content(content, output) - - except Exception as e: - print(f"\nError during API call: {e}", file=sys.stderr) - - -def _process_response_content(content: str, output_base: Optional[str]): - """ - Extract logic to find Python code, save it, and attempt to run it to generate JSON. - """ - # Look for python code block - code_match = re.search(r"```python\n(.*?)\n```", content, re.DOTALL) - if not code_match: - return - - code = code_match.group(1) - - # Determine output filenames - if output_base: - py_file = Path(output_base + ".py") - json_file = Path(output_base + ".json") - else: - # Default name - py_file = Path("cohort_definition.py") - json_file = Path("cohort_definition.json") - - # Save Python code - try: - py_file.write_text(code) - print(f"\n✅ Saved Python code to {py_file}") - except Exception as e: - print(f"Error saving Python file: {e}") - return - - # Attempt to execute and save JSON - # This involves running the code and capturing the 'cohort' variable or 'expression' variable - print(" Attempting to generate JSON...") - - try: - # Create a local scope - local_scope = {} - # We need to make sure the CWD is in path so imports work? - # Assuming we are running from project root or installed package - - exec(code, {}, local_scope) - - # Look for a CohortExpression or CohortBuilder object - # The prompt usually produces: - # cohort = CohortBuilder(...).build() - # So we look for 'cohort' - - cohort_obj = local_scope.get("cohort") - if not cohort_obj: - # Try to find any variable that is a tuple (builder) or CohortExpression - for _k, v in local_scope.items(): - if hasattr(v, "to_json"): # CohortExpression has to_json? Check API. - cohort_obj = v - break - - if cohort_obj: - # If it's the builder (tuple in some cases?), checks if it has build() - # But the prompt says `.build()` returns CohortExpression. - - # Check if it has 'to_json' or similar. - # circe.cohortdefinition.CohortExpression uses Pydantic? - # It inherits from Serializable? - - json_output = None - if hasattr(cohort_obj, "json"): # Pydantic v1/v2 - json_output = ( - cohort_obj.model_dump_json(indent=2) - if hasattr(cohort_obj, "model_dump_json") - else cohort_obj.json(indent=2) - ) - elif hasattr(cohort_obj, "to_json"): - json_output = cohort_obj.to_json() - else: - # It might be a dict? - if isinstance(cohort_obj, dict): - json_output = json.dumps(cohort_obj, indent=2) - - if json_output: - json_file.write_text(json_output) - print(f"✅ Saved Cohort JSON to {json_file}") - else: - print(" Could not serialize 'cohort' object to JSON.") - else: - print(" Could not find 'cohort' variable in executed code.") - - except Exception as e: - print(f" Error executing generated code: {e}") - print(" (Ensure the generated code is valid and all dependencies are installed)") diff --git a/circe/check/checkers/attribute_checker_factory.py b/circe/check/checkers/attribute_checker_factory.py index 2d7681de..7b918c36 100644 --- a/circe/check/checkers/attribute_checker_factory.py +++ b/circe/check/checkers/attribute_checker_factory.py @@ -8,7 +8,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Any, Callable +from collections.abc import Callable +from typing import Any from ..constants import Constants from .base_checker_factory import BaseCheckerFactory diff --git a/circe/check/checkers/base_checker_factory.py b/circe/check/checkers/base_checker_factory.py index 10b3eb05..78839ba5 100644 --- a/circe/check/checkers/base_checker_factory.py +++ b/circe/check/checkers/base_checker_factory.py @@ -9,7 +9,7 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable +from collections.abc import Callable from .warning_reporter import WarningReporter diff --git a/circe/check/checkers/comparisons.py b/circe/check/checkers/comparisons.py index 29306aa2..8c379cca 100644 --- a/circe/check/checkers/comparisons.py +++ b/circe/check/checkers/comparisons.py @@ -76,7 +76,7 @@ def start_is_greater_than_end(range_val) -> bool: return False @staticmethod - def is_date_valid(date: Optional[str]) -> bool: + def is_date_valid(date: str | None) -> bool: """Check if a date string is valid. Args: diff --git a/circe/check/checkers/concept_checker_factory.py b/circe/check/checkers/concept_checker_factory.py index 118058d5..c987cc7a 100644 --- a/circe/check/checkers/concept_checker_factory.py +++ b/circe/check/checkers/concept_checker_factory.py @@ -8,7 +8,7 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable, Optional +from collections.abc import Callable from ..constants import Constants from ..operations.operations import Operations @@ -427,7 +427,7 @@ def check(c: "DemographicCriteria") -> None: return check - def _check_concept(self, concepts: Optional[list["Concept"]], criteria_name: str, attribute: str) -> None: + def _check_concept(self, concepts: list["Concept"] | None, criteria_name: str, attribute: str) -> None: """Check if a concept array is empty. Args: diff --git a/circe/check/checkers/concept_set_selection_checker_factory.py b/circe/check/checkers/concept_set_selection_checker_factory.py index 3967a9ff..3e08370f 100644 --- a/circe/check/checkers/concept_set_selection_checker_factory.py +++ b/circe/check/checkers/concept_set_selection_checker_factory.py @@ -8,7 +8,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable, Optional +from collections.abc import Callable +from typing import Optional from ..constants import Constants from ..operations.operations import Operations diff --git a/circe/check/checkers/criteria_checker_factory.py b/circe/check/checkers/criteria_checker_factory.py index 47f19fe0..69581219 100644 --- a/circe/check/checkers/criteria_checker_factory.py +++ b/circe/check/checkers/criteria_checker_factory.py @@ -8,7 +8,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable, Optional +from collections.abc import Callable +from typing import Optional # Import at runtime to avoid circular dependencies try: @@ -213,7 +214,7 @@ def _get_concept_set_selection_suppliers( Returns: A list of functions that return ConceptSetSelection objects """ - suppliers: list[Callable[[], Optional[ConceptSetSelection]]] = [] + suppliers: list[Callable[[], ConceptSetSelection | None]] = [] suppliers.append(lambda: criteria.place_of_service_cs) suppliers.append(lambda: criteria.gender_cs) suppliers.append(lambda: criteria.provider_specialty_cs) diff --git a/circe/check/checkers/drug_domain_check.py b/circe/check/checkers/drug_domain_check.py index c251a726..3d14e37f 100644 --- a/circe/check/checkers/drug_domain_check.py +++ b/circe/check/checkers/drug_domain_check.py @@ -78,7 +78,7 @@ def _check(self, expression: "CohortExpression", reporter: WarningReporter) -> N title = "Concept sets" if len(concept_sets) > 1 else "Concept set" reporter(self.MESSAGE, title, names) - def _map_criteria(self, criteria: "Criteria") -> Optional[int]: + def _map_criteria(self, criteria: "Criteria") -> int | None: """Map a criteria to its codeset ID. Args: diff --git a/circe/check/checkers/events_progression_check.py b/circe/check/checkers/events_progression_check.py index 5a46211d..e7c04e75 100644 --- a/circe/check/checkers/events_progression_check.py +++ b/circe/check/checkers/events_progression_check.py @@ -38,7 +38,7 @@ class LimitType(Enum): LATEST = (1, "Last") ALL = (2, "All") - def __init__(self, weight: int, name: Optional[str]): + def __init__(self, weight: int, name: str | None): """Initialize a limit type. Args: @@ -58,7 +58,7 @@ def weight(self) -> int: return self._weight @property - def name(self) -> Optional[str]: + def name(self) -> str | None: """Get the name of this limit type. Returns: @@ -67,7 +67,7 @@ def name(self) -> Optional[str]: return self._name @staticmethod - def from_name(name: Optional[str]) -> "LimitType": + def from_name(name: str | None) -> "LimitType": """Get a limit type from its name. Args: diff --git a/circe/check/checkers/range_checker_factory.py b/circe/check/checkers/range_checker_factory.py index 3c9f36cd..a9e5ab10 100644 --- a/circe/check/checkers/range_checker_factory.py +++ b/circe/check/checkers/range_checker_factory.py @@ -8,7 +8,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Any, Callable, Optional +from collections.abc import Callable +from typing import Any, Optional from ..constants import Constants from ..operations.operations import Operations diff --git a/circe/check/checkers/text_checker_factory.py b/circe/check/checkers/text_checker_factory.py index d002f2da..42f1a82e 100644 --- a/circe/check/checkers/text_checker_factory.py +++ b/circe/check/checkers/text_checker_factory.py @@ -8,7 +8,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable, Optional +from collections.abc import Callable +from typing import Optional from ..constants import Constants from ..operations.operations import Operations diff --git a/circe/check/checkers/time_window_check.py b/circe/check/checkers/time_window_check.py index 410f34d1..e9aefbe2 100644 --- a/circe/check/checkers/time_window_check.py +++ b/circe/check/checkers/time_window_check.py @@ -8,7 +8,7 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Any, Optional +from typing import Any from ..operations.operations import Operations from ..utils.criteria_name_helper import CriteriaNameHelper @@ -42,7 +42,7 @@ class TimeWindowCheck(BaseCorelatedCriteriaCheck): def __init__(self): """Initialize the time window check.""" super().__init__() - self._observation_filter: Optional[ObservationFilter] = None + self._observation_filter: ObservationFilter | None = None def _define_severity(self) -> WarningSeverity: """Define the severity level for this check. diff --git a/circe/check/checkers/unused_concepts_check.py b/circe/check/checkers/unused_concepts_check.py index 3cf065a5..b0da7e8e 100644 --- a/circe/check/checkers/unused_concepts_check.py +++ b/circe/check/checkers/unused_concepts_check.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..warning_severity import WarningSeverity from ..warnings.concept_set_warning import ConceptSetWarning from .base_check import BaseCheck @@ -243,7 +241,7 @@ def _correlated_criteria_to_list(self, correlated_criteria) -> list["Criteria"]: ) return criteria_list - def _to_criteria_list(self, criteria_list: Optional[list["CorelatedCriteria"]]) -> list["Criteria"]: + def _to_criteria_list(self, criteria_list: list["CorelatedCriteria"] | None) -> list["Criteria"]: """Convert a list of CorelatedCriteria to a list of Criteria. Args: @@ -256,7 +254,7 @@ def _to_criteria_list(self, criteria_list: Optional[list["CorelatedCriteria"]]) return [] return [c.criteria for c in criteria_list if hasattr(c, "criteria") and c.criteria] - def _to_criteria_list_from_groups(self, groups: Optional[list["CriteriaGroup"]]) -> list["Criteria"]: + def _to_criteria_list_from_groups(self, groups: list["CriteriaGroup"] | None) -> list["Criteria"]: """Convert groups to a list of criteria. Args: diff --git a/circe/check/operations/__init__.py b/circe/check/operations/__init__.py index b23bc8c6..3b6fb8e2 100644 --- a/circe/check/operations/__init__.py +++ b/circe/check/operations/__init__.py @@ -5,7 +5,7 @@ """ # Type alias for convenience (Callable[[], None]) -from typing import Callable +from collections.abc import Callable from .conditional_operations import ConditionalOperations from .execution import Execution diff --git a/circe/check/operations/conditional_operations.py b/circe/check/operations/conditional_operations.py index 700aadbd..ed38f625 100644 --- a/circe/check/operations/conditional_operations.py +++ b/circe/check/operations/conditional_operations.py @@ -9,7 +9,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import TYPE_CHECKING, Callable, Generic, Protocol, TypeVar +from collections.abc import Callable +from typing import TYPE_CHECKING, Generic, Protocol, TypeVar if TYPE_CHECKING: from .executive_operations import ExecutiveOperations diff --git a/circe/check/operations/executive_operations.py b/circe/check/operations/executive_operations.py index aa1ce75a..1b0b266d 100644 --- a/circe/check/operations/executive_operations.py +++ b/circe/check/operations/executive_operations.py @@ -9,7 +9,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Callable, Generic, Protocol, TypeVar, overload +from collections.abc import Callable +from typing import Generic, Protocol, TypeVar, overload from .conditional_operations import ConditionalOperations from .execution import Execution diff --git a/circe/check/operations/operations.py b/circe/check/operations/operations.py index 3e4f0ca6..0faa1993 100644 --- a/circe/check/operations/operations.py +++ b/circe/check/operations/operations.py @@ -9,7 +9,8 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Any, Callable, Generic, Optional, TypeVar +from collections.abc import Callable +from typing import Any, Generic, TypeVar from .conditional_operations import ConditionalOperations from .executive_operations import ExecutiveOperations @@ -34,8 +35,8 @@ def __init__(self, value: T): value: The value to match against """ self._value = value - self._result: Optional[bool] = None - self._return_value: Optional[V] = None + self._result: bool | None = None + self._return_value: V | None = None @staticmethod def match(value: T) -> ConditionalOperations[T, V]: @@ -113,7 +114,7 @@ def or_else(self, consumer: Callable[[T], None]) -> None: if not self._result: consumer(self._value) - def value(self) -> Optional[V]: + def value(self) -> V | None: """Get the return value from then_return operations. Returns: diff --git a/circe/check/warnings/concept_set_warning.py b/circe/check/warnings/concept_set_warning.py index 0be3c8c2..d7cdc878 100644 --- a/circe/check/warnings/concept_set_warning.py +++ b/circe/check/warnings/concept_set_warning.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ...vocabulary.concept import ConceptSet from ..warning_severity import WarningSeverity from .base_warning import BaseWarning @@ -28,7 +26,7 @@ def __init__( self, severity: WarningSeverity, template: str, - concept_set: Optional[ConceptSet], + concept_set: ConceptSet | None, ): """Initialize a concept set warning. @@ -42,7 +40,7 @@ def __init__( self._concept_set = concept_set @property - def concept_set(self) -> Optional[ConceptSet]: + def concept_set(self) -> ConceptSet | None: """Get the concept set associated with this warning. Returns: diff --git a/circe/cohort_definition_set/__init__.py b/circe/cohort_definition_set/__init__.py new file mode 100644 index 00000000..14bc7550 --- /dev/null +++ b/circe/cohort_definition_set/__init__.py @@ -0,0 +1,35 @@ +"""CohortDefinitionSet — batch cohort generation with incremental caching. + +This module provides the Python equivalent of OHDSI/CohortGenerator's +CohortDefinitionSet: a typed container for multiple cohort definitions that can +be generated simultaneously against an ibis backend, with optional checksum-based +incremental skipping. + +Example: + >>> from circe.cohort_definition_set import ( + ... CohortDefinitionSet, + ... generate_cohort_set, + ... ) + >>> cds = CohortDefinitionSet() + >>> cds.add(cohort_id=1, cohort_name="Diabetes", expression=expr1) + >>> cds.add(cohort_id=2, cohort_name="Hypertension", expression=expr2) + >>> results = generate_cohort_set( + ... cds, + ... backend=conn, + ... cdm_schema="main", + ... cohort_table="cohort", + ... incremental=True, + ... ) +""" + +from ._core import CohortDefinition, CohortDefinitionSet, CohortGenerationResult +from ._generate import async_generate_cohort_set, generate_cohort_set, summarise_generation_results + +__all__ = [ + "CohortDefinition", + "CohortDefinitionSet", + "CohortGenerationResult", + "async_generate_cohort_set", + "generate_cohort_set", + "summarise_generation_results", +] diff --git a/circe/cohort_definition_set/_checksum_store.py b/circe/cohort_definition_set/_checksum_store.py new file mode 100644 index 00000000..21b0d551 --- /dev/null +++ b/circe/cohort_definition_set/_checksum_store.py @@ -0,0 +1,289 @@ +"""Persistent generation history for incremental cohort generation. + +The generation history table records the SHA-256 checksum of each generated +cohort's expression alongside its generation status, start time, and end time. +On subsequent incremental runs, cohorts whose expression checksum matches the +most recent stored value are skipped. + +This table serves as the canonical source of truth for per-cohort generation +timing, enabling fair benchmarks across implementations. + +Table schema (v2, introduced 0.3.0): + cohort_definition_id int64 + checksum str + status str -- "COMPLETE" or "FAILED" + start_time timestamp + end_time timestamp + +The original v1 schema stored only ``(cohort_definition_id, checksum, +generation_end_time)`` and is still handled transparently for reads. +""" + +from __future__ import annotations + +from datetime import datetime +from typing import TYPE_CHECKING + +from ..execution.ibis.operations import create_table, read_table, table_exists + +if TYPE_CHECKING: + from ..execution.typing import IbisBackendLike, Table + + +def load_checksums( + backend: IbisBackendLike, + *, + schema: str | None, + table_name: str, +) -> dict[int, str]: + """Load stored checksums from the generation history table. + + Returns a mapping of cohort_id -> checksum for the most recently recorded + completed generation of each cohort. Returns an empty dict if the table + does not yet exist. + + Args: + backend: Ibis backend connection. + schema: Schema/database where the table lives. + table_name: Name of the table (may be v1 ``cohort_checksum`` or v2 + ``cohort_generation_history`` format). + + Returns: + dict mapping cohort_id (int) -> checksum (str). + """ + if not table_exists(backend, table_name=table_name, schema=schema): + return {} + + import ibis + + table = read_table(backend, table_name=table_name, schema=schema) + column_names = table.schema().names + has_status = "status" in column_names + has_end_time = "end_time" in column_names + has_gen_end_time = "generation_end_time" in column_names + time_col = "end_time" if has_end_time else ("generation_end_time" if has_gen_end_time else None) + + if has_status: + table = table.filter(table.status == ibis.literal("COMPLETE", type="str")) + + if time_col is not None: + w = ibis.window( + group_by=table.cohort_definition_id, + order_by=ibis.desc(table[time_col]), + ) + ranked = table.mutate(_rn=ibis.row_number().over(w)) + table = ranked.filter(ranked._rn == 0) + + rows = table.select("cohort_definition_id", "checksum").execute() + if rows.empty: + return {} + return {int(row["cohort_definition_id"]): str(row["checksum"]) for _, row in rows.iterrows()} + + +def load_generation_history( + backend: IbisBackendLike, + *, + schema: str | None, + table_name: str, +) -> Table | None: + """Load the full generation history from the history table. + + Returns an ibis Table expression with all columns (cohort_definition_id, + checksum, status, start_time, end_time) for every recorded generation. + Returns ``None`` if the table does not exist or was created with the v1 + schema that lacks timing columns. + + Args: + backend: Ibis backend connection. + schema: Schema/database where the table lives. + table_name: Name of the generation history table. + + Returns: + ibis Table with per-cohort history, or ``None`` if unavailable. + """ + if not table_exists(backend, table_name=table_name, schema=schema): + return None + + table = read_table(backend, table_name=table_name, schema=schema) + column_names = table.schema().names + if "start_time" not in column_names or "status" not in column_names: + return None + + return table + + +def save_checksums( + backend: IbisBackendLike, + *, + schema: str | None, + table_name: str, + completed: dict[int, tuple[str, datetime]], +) -> None: + """Persist checksums for successfully generated cohorts (v1 compat). + + .. deprecated:: 0.3.0 + Prefer ``save_generation_history()`` which stores the full generation + record including status and start_time. This wrapper is retained for + backward compatibility and delegates internally. + + Args: + backend: Ibis backend connection. + schema: Schema/database where the table should be written. + table_name: Name of the table. + completed: Mapping of cohort_id -> (checksum, end_time) for cohorts + that completed successfully in this run. + """ + if not completed: + return + + now = datetime.now() + converted: dict[int, tuple[str, str, datetime, datetime]] = {} + for cohort_id, (checksum, end_time) in completed.items(): + converted[cohort_id] = (checksum, "COMPLETE", now, end_time) + + save_generation_history( + backend, + schema=schema, + table_name=table_name, + generated=converted, + ) + + +def save_generation_history( + backend: IbisBackendLike, + *, + schema: str | None, + table_name: str, + generated: dict[int, tuple[str, str, datetime, datetime]], +) -> None: + """Persist generation history for all generated cohorts (COMPLETE and FAILED). + + Uses the same read-filter-union-rewrite pattern as ``write_cohort`` so it + works on every ibis backend without requiring raw SQL. + + Each row written to the table contains: ``(cohort_definition_id, checksum, + status, start_time, end_time)``. SKIPPED cohorts are intentionally + omitted so their prior history entry is preserved. + + Args: + backend: Ibis backend connection. + schema: Schema/database where the table should be written. + table_name: Name of the generation history table. + generated: Mapping of cohort_id -> (checksum, status, start_time, + end_time) for every cohort that was processed in this run + (COMPLETE or FAILED). + """ + if not generated: + return + + import ibis + + def _checksum_row(cid, checksum, status, start_time, end_time): + return ( + ibis.literal(int(cid), type="int64") + .name("cohort_definition_id") + .as_table() + .mutate( + checksum=ibis.literal(str(checksum), type="str"), + status=ibis.literal(str(status), type="str"), + start_time=ibis.literal(start_time, type="timestamp"), + end_time=ibis.literal(end_time, type="timestamp"), + ) + ) + + items = list(generated.items()) + new_relation = _checksum_row(items[0][0], *items[0][1]) + for cid, vals in items[1:]: + new_relation = new_relation.union(_checksum_row(cid, *vals), distinct=False) + + if not table_exists(backend, table_name=table_name, schema=schema): + create_table(backend, table_name=table_name, schema=schema, obj=new_relation, overwrite=False) + return + + existing = read_table(backend, table_name=table_name, schema=schema) + updated_ids = list(generated.keys()) + filtered_existing = existing.filter( + ~existing.cohort_definition_id.cast("int64").isin( + [ibis.literal(int(i), type="int64") for i in updated_ids] + ) + ) + end_ts_type = existing.schema()["end_time"] + start_ts_type = existing.schema()["start_time"] + new_relation = new_relation.mutate( + start_time=new_relation.start_time.cast(start_ts_type), + end_time=new_relation.end_time.cast(end_ts_type), + ) + merged = filtered_existing.union(new_relation, distinct=False) + create_table(backend, table_name=table_name, schema=schema, obj=merged, overwrite=True) + + +def upsert_generation_history( + backend: IbisBackendLike, + *, + schema: str | None, + table_name: str, + cohort_id: int, + checksum: str, + status: str, + start_time: datetime, + end_time: datetime, +) -> None: + """Persist the generation result for a single cohort. + + Unlike ``save_generation_history()`` this operates on one cohort at a + time so that incremental persistence is possible — a completed cohort + is recorded immediately rather than waiting for the entire batch to + finish. + + Uses DELETE + INSERT (no full-table rewrite) so it is O(1) per call. + + Args: + backend: Ibis backend connection. + schema: Schema/database where the table lives. + table_name: Name of the generation history table. + cohort_id: Cohort definition id. + checksum: Expression checksum for this generation. + status: ``"COMPLETE"`` or ``"FAILED"``. + start_time: When execution started. + end_time: When execution ended. + """ + import ibis + + from ..execution.ibis.operations import ( + create_table, + delete_cohort_rows, + insert_rows_via_raw_sql, + table_exists, + ) + + columns = ["cohort_definition_id", "checksum", "status", "start_time", "end_time"] + row = [int(cohort_id), str(checksum), str(status), start_time, end_time] + + if not table_exists(backend, table_name=table_name, schema=schema): + new_row = ( + ibis.literal(int(cohort_id), type="int64") + .name("cohort_definition_id") + .as_table() + .mutate( + checksum=ibis.literal(str(checksum), type="str"), + status=ibis.literal(str(status), type="str"), + start_time=ibis.literal(start_time, type="timestamp"), + end_time=ibis.literal(end_time, type="timestamp"), + ) + ) + create_table(backend, table_name=table_name, schema=schema, obj=new_row, overwrite=False) + return + + delete_cohort_rows( + backend, + cohort_table=table_name, + results_schema=schema, + cohort_id=cohort_id, + ) + insert_rows_via_raw_sql( + backend, + table_name=table_name, + schema=schema, + columns=columns, + rows=[row], + ) diff --git a/circe/cohort_definition_set/_core.py b/circe/cohort_definition_set/_core.py new file mode 100644 index 00000000..ddc4081c --- /dev/null +++ b/circe/cohort_definition_set/_core.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +from collections.abc import Iterator +from dataclasses import dataclass, field +from datetime import datetime +from typing import TYPE_CHECKING, Literal + +if TYPE_CHECKING: + from ..cohortdefinition.cohort import CohortExpression +else: + from ..cohortdefinition.cohort import CohortExpression + + +@dataclass +class CohortDefinition: + """A single cohort entry in a CohortDefinitionSet.""" + + cohort_id: int + cohort_name: str + expression: CohortExpression + + +@dataclass +class CohortGenerationResult: + """Result of generating a single cohort from a CohortDefinitionSet.""" + + cohort_id: int + cohort_name: str + status: Literal["COMPLETE", "SKIPPED", "FAILED"] + checksum: str + start_time: datetime + end_time: datetime + error: Exception | None = field(default=None, compare=False) + + +class CohortDefinitionSet: + """Container for a collection of cohort definitions to be generated together. + + Modelled after OHDSI/CohortGenerator's CohortDefinitionSet, but using typed + Python classes rather than an R data.frame. + + Example: + >>> cds = CohortDefinitionSet() + >>> cds.add(cohort_id=1, cohort_name="Diabetes", expression=expr1) + >>> cds.add(cohort_id=2, cohort_name="Hypertension", expression=expr2) + >>> len(cds) + 2 + """ + + def __init__(self) -> None: + self._cohorts: list[CohortDefinition] = [] + self._id_index: dict[int, int] = {} # cohort_id -> list index + + def add(self, cohort_id: int, cohort_name: str, expression: CohortExpression) -> None: + """Add a cohort definition to this set. + + Args: + cohort_id: Unique integer identifier for this cohort. + cohort_name: Human-readable name for this cohort. + expression: The CohortExpression defining the cohort logic. + + Raises: + ValueError: If a cohort with the same cohort_id already exists. + """ + if cohort_id in self._id_index: + raise ValueError( + f"A cohort with cohort_id={cohort_id} already exists in this CohortDefinitionSet." + ) + self._id_index[cohort_id] = len(self._cohorts) + self._cohorts.append( + CohortDefinition(cohort_id=cohort_id, cohort_name=cohort_name, expression=expression) + ) + + def __len__(self) -> int: + return len(self._cohorts) + + def __iter__(self) -> Iterator[CohortDefinition]: + return iter(self._cohorts) + + def __getitem__(self, cohort_id: int) -> CohortDefinition: + """Retrieve a cohort definition by its cohort_id. + + Raises: + KeyError: If no cohort with the given id exists. + """ + if cohort_id not in self._id_index: + raise KeyError(f"No cohort with cohort_id={cohort_id} in this CohortDefinitionSet.") + return self._cohorts[self._id_index[cohort_id]] + + def checksums(self) -> dict[int, str]: + """Return a mapping of cohort_id to the expression checksum for each cohort. + + Checksums are computed using CohortExpression.checksum(), which normalises + the expression JSON and produces a SHA-256 hex digest. This is suitable for + detecting whether a cohort definition has changed between runs. + + Returns: + dict mapping cohort_id -> hex checksum string + """ + return {c.cohort_id: c.expression.checksum() for c in self._cohorts} diff --git a/circe/cohort_definition_set/_generate.py b/circe/cohort_definition_set/_generate.py new file mode 100644 index 00000000..ce981ddb --- /dev/null +++ b/circe/cohort_definition_set/_generate.py @@ -0,0 +1,379 @@ +"""Batch cohort generation for CohortDefinitionSet.""" + +from __future__ import annotations + +import asyncio +import contextlib +import logging +import threading +import uuid +from datetime import datetime +from typing import TYPE_CHECKING, Literal + +from ..execution.api import build_cohort, write_cohort +from ..execution.ibis.materialize import project_to_ohdsi_cohort_table +from ._checksum_store import load_checksums, upsert_generation_history +from ._core import CohortDefinition, CohortDefinitionSet, CohortGenerationResult + +if TYPE_CHECKING: + from ..execution.typing import IbisBackendLike + +logger = logging.getLogger(__name__) + +_backend_lock = threading.Lock() + + +def _drop_tables_by_prefix(backend: IbisBackendLike, prefix: str, schema: str | None) -> None: + """Drop all tables in *schema* whose name starts with *prefix*.""" + try: + tables = backend.list_tables(database=schema) + except Exception: + tables = backend.list_tables() + for table_name in tables: + if table_name.startswith(prefix): + with contextlib.suppress(Exception): + backend.drop_table(table_name, database=schema, force=True) + + +def _process_single_cohort( + cohort: CohortDefinition, + *, + backend: IbisBackendLike, + cdm_schema: str | None, + results_schema: str | None, + vocabulary_schema: str | None, + cohort_table: str, + session_prefix: str, +) -> tuple[datetime, datetime]: + """Build and write a single cohort. Thread-safe via ``_backend_lock``. + + Each cohort gets its own per-cohort codeset table populated and + dropped as it runs, mirroring the Java ``#Codesets`` pattern. + """ + with _backend_lock: + start_time = datetime.now() + new_rows = build_cohort( + cohort.expression, + backend=backend, + cdm_schema=cdm_schema, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + cohort_id=cohort.cohort_id, + cohort_table=cohort_table, + session_prefix=session_prefix, + ) + projected = project_to_ohdsi_cohort_table(new_rows, cohort_id=cohort.cohort_id) + write_cohort( + compiled_relation=projected, + backend=backend, + cdm_schema=cdm_schema, + cohort_table=cohort_table, + cohort_id=cohort.cohort_id, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + if_exists="replace", + ) + end_time = datetime.now() + return start_time, end_time + + +async def async_generate_cohort_set( + cohort_definition_set: CohortDefinitionSet, + *, + backend: IbisBackendLike, + cdm_schema: str, + cohort_table: str, + results_schema: str | None = None, + vocabulary_schema: str | None = None, + incremental: bool = False, + checksum_table: str | None = None, + stop_on_error: bool = True, + compile_timeout: float | None = None, +) -> list[CohortGenerationResult]: + """Generate all cohorts in a CohortDefinitionSet and write them to a shared table. + + Each cohort builds its own per-cohort codeset table. When *incremental* + is True, concept sets are stored in the persistent + ``_circe_codeset_cache`` keyed by SHA-256 checksum so that identical + concept set definitions across cohorts are resolved only once per + batch. + + All exception types are caught and recorded as ``FAILED``. + + Args: + cohort_definition_set: The set of cohort definitions to generate. + backend: Ibis backend connection pointing at the target database. + cdm_schema: Schema containing the OMOP CDM source tables. + cohort_table: Name of the OHDSI cohort table to write results into. + results_schema: Optional schema for both the cohort table and + checksum table. + vocabulary_schema: Optional schema for vocabulary tables. + incremental: If True, skip cohorts whose expression checksum is + unchanged since the last successful generation. + checksum_table: Name of the table used to persist generation + history for incremental runs. + stop_on_error: If True, raise on the first failure. + compile_timeout: Maximum seconds per-cohort before timeout. + + Returns: + A list of :class:`CohortGenerationResult`. + """ + total = len(cohort_definition_set) + + from ..execution.engine.group_operators import _COMPILED_CORRELATED_EVENTS + + _COMPILED_CORRELATED_EVENTS.clear() + + session_prefix = f"__s_{uuid.uuid4().hex[:8]}_" + + if checksum_table is None: + checksum_table = f"{cohort_table}_checksum" + + previous_checksums: dict[int, str] = {} + if incremental: + previous_checksums = await asyncio.to_thread( + load_checksums, + backend, + schema=results_schema, + table_name=checksum_table, + ) + + results: list[CohortGenerationResult] = [] + + logger.info("Generating %d cohort(s) (incremental=%s)", total, incremental) + + for i, cohort in enumerate(cohort_definition_set, start=1): + current_checksum = cohort.expression.checksum() + + if incremental and previous_checksums.get(cohort.cohort_id) == current_checksum: + logger.info( + "[%d/%d] Skipping cohort %d (%s) -- checksum unchanged", + i, + total, + cohort.cohort_id, + cohort.cohort_name, + ) + results.append( + CohortGenerationResult( + cohort_id=cohort.cohort_id, + cohort_name=cohort.cohort_name, + status="SKIPPED", + checksum=current_checksum, + start_time=datetime.now(), + end_time=datetime.now(), + ) + ) + continue + + logger.info( + "[%d/%d] Building cohort %d (%s) ...", + i, + total, + cohort.cohort_id, + cohort.cohort_name, + ) + + start_time: datetime | None = None + end_time: datetime | None = None + try: + start_time, end_time = await asyncio.wait_for( + asyncio.to_thread( + _process_single_cohort, + cohort, + backend=backend, + cdm_schema=cdm_schema, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + cohort_table=cohort_table, + session_prefix=session_prefix, + ), + timeout=compile_timeout, + ) + + duration = (end_time - start_time).total_seconds() + logger.info( + "[%d/%d] Completed cohort %d (%s) -- duration %.1fs", + i, + total, + cohort.cohort_id, + cohort.cohort_name, + duration, + ) + except asyncio.TimeoutError: + if end_time is None: + end_time = datetime.now() + duration = (end_time - (start_time or end_time)).total_seconds() + logger.error( + "[%d/%d] TIMED OUT cohort %d (%s) after %.1fs", + i, + total, + cohort.cohort_id, + cohort.cohort_name, + duration, + ) + timeout_exc = TimeoutError( + f"Cohort {cohort.cohort_id} ({cohort.cohort_name}) exceeded timeout of {compile_timeout:.0f}s" + ) + results.append( + CohortGenerationResult( + cohort_id=cohort.cohort_id, + cohort_name=cohort.cohort_name, + status="FAILED", + checksum=current_checksum, + start_time=start_time or datetime.now(), + end_time=end_time, + error=timeout_exc, + ) + ) + if incremental: + upsert_generation_history( + backend, + schema=results_schema, + table_name=checksum_table, + cohort_id=cohort.cohort_id, + checksum=current_checksum, + status="FAILED", + start_time=start_time or datetime.now(), + end_time=end_time, + ) + if stop_on_error: + raise timeout_exc from None + continue + except Exception as exc: + if end_time is None: + end_time = datetime.now() + duration = (end_time - (start_time or end_time)).total_seconds() + logger.error( + "[%d/%d] FAILED cohort %d (%s) after %.1fs: %s", + i, + total, + cohort.cohort_id, + cohort.cohort_name, + duration, + exc, + ) + results.append( + CohortGenerationResult( + cohort_id=cohort.cohort_id, + cohort_name=cohort.cohort_name, + status="FAILED", + checksum=current_checksum, + start_time=start_time or datetime.now(), + end_time=end_time, + error=exc, + ) + ) + if incremental: + upsert_generation_history( + backend, + schema=results_schema, + table_name=checksum_table, + cohort_id=cohort.cohort_id, + checksum=current_checksum, + status="FAILED", + start_time=start_time or datetime.now(), + end_time=end_time, + ) + if stop_on_error: + raise + continue + + # Individual staging tables are cleaned by prefix at batch end + + results.append( + CohortGenerationResult( + cohort_id=cohort.cohort_id, + cohort_name=cohort.cohort_name, + status="COMPLETE", + checksum=current_checksum, + start_time=start_time or datetime.now(), + end_time=end_time or datetime.now(), + ) + ) + if incremental: + upsert_generation_history( + backend, + schema=results_schema, + table_name=checksum_table, + cohort_id=cohort.cohort_id, + checksum=current_checksum, + status="COMPLETE", + start_time=start_time or datetime.now(), + end_time=end_time or datetime.now(), + ) + + summary = summarise_generation_results(results) + logger.info( + "Cohort generation complete: %d completed, %d skipped, %d failed", + summary["COMPLETE"], + summary["SKIPPED"], + summary["FAILED"], + ) + + # Drop all staging tables from this batch run + await asyncio.to_thread( + _drop_tables_by_prefix, + backend, + session_prefix, + results_schema or cdm_schema, + ) + + return results + + +def generate_cohort_set( + cohort_definition_set: CohortDefinitionSet, + *, + backend: IbisBackendLike, + cdm_schema: str, + cohort_table: str, + results_schema: str | None = None, + vocabulary_schema: str | None = None, + incremental: bool = False, + checksum_table: str | None = None, + stop_on_error: bool = True, +) -> list[CohortGenerationResult]: + """Generate all cohorts in a CohortDefinitionSet and write them to a shared table. + + This synchronous wrapper delegates to :func:`async_generate_cohort_set` + via :func:`asyncio.run`. See that function for full parameter + documentation. + + Raises: + RuntimeError: If called from within a running asyncio event loop. + Use :func:`async_generate_cohort_set` directly in that case. + """ + return asyncio.run( + async_generate_cohort_set( + cohort_definition_set, + backend=backend, + cdm_schema=cdm_schema, + cohort_table=cohort_table, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + incremental=incremental, + checksum_table=checksum_table, + stop_on_error=stop_on_error, + ) + ) + + +def summarise_generation_results( + results: list[CohortGenerationResult], +) -> dict[Literal["COMPLETE", "SKIPPED", "FAILED"], int]: + """Return a count summary of generation results by status. + + Args: + results: List of CohortGenerationResult from generate_cohort_set. + + Returns: + dict with counts for each status. + """ + counts: dict[Literal["COMPLETE", "SKIPPED", "FAILED"], int] = { + "COMPLETE": 0, + "SKIPPED": 0, + "FAILED": 0, + } + for r in results: + counts[r.status] += 1 + return counts diff --git a/circe/cohortdefinition/builders/base.py b/circe/cohortdefinition/builders/base.py index 3cf9c9f4..4d158ba9 100644 --- a/circe/cohortdefinition/builders/base.py +++ b/circe/cohortdefinition/builders/base.py @@ -10,7 +10,7 @@ """ from abc import ABC, abstractmethod -from typing import Generic, Optional, TypeVar +from typing import Generic, TypeVar from ..criteria import Criteria from .utils import BuilderOptions, CriteriaColumn @@ -24,14 +24,14 @@ class CriteriaSqlBuilder(ABC, Generic[T]): Java equivalent: org.ohdsi.circe.cohortdefinition.builders.CriteriaSqlBuilder """ - def get_criteria_sql(self, criteria: T, options: Optional[BuilderOptions] = None) -> str: + def get_criteria_sql(self, criteria: T, options: BuilderOptions | None = None) -> str: """Get SQL query for criteria. Java equivalent: CriteriaSqlBuilder.getCriteriaSql(T criteria) """ return self.get_criteria_sql_with_options(criteria, options) - def get_criteria_sql_with_options(self, criteria: T, options: Optional[BuilderOptions]) -> str: + def get_criteria_sql_with_options(self, criteria: T, options: BuilderOptions | None) -> str: """Get SQL query for criteria with builder options. Java equivalent: CriteriaSqlBuilder.getCriteriaSql(T criteria, BuilderOptions options) @@ -99,7 +99,7 @@ def embed_codeset_clause(self, query: str, criteria: T) -> str: # This would need to be implemented based on the Java logic return query.replace("@codesetClause", "") - def resolve_select_clauses(self, criteria: T, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_select_clauses(self, criteria: T, options: BuilderOptions | None = None) -> list[str]: """Resolve select clauses for criteria. Java equivalent: CriteriaSqlBuilder.resolveSelectClauses() @@ -107,7 +107,7 @@ def resolve_select_clauses(self, criteria: T, options: Optional[BuilderOptions] # This would need to be implemented based on the Java logic return [] - def resolve_join_clauses(self, criteria: T, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: T, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for criteria. Java equivalent: CriteriaSqlBuilder.resolveJoinClauses() @@ -115,7 +115,7 @@ def resolve_join_clauses(self, criteria: T, options: Optional[BuilderOptions] = # This would need to be implemented based on the Java logic return [] - def resolve_where_clauses(self, criteria: T, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_where_clauses(self, criteria: T, options: BuilderOptions | None = None) -> list[str]: """Resolve where clauses for criteria. Java equivalent: CriteriaSqlBuilder.resolveWhereClauses() diff --git a/circe/cohortdefinition/builders/condition_era.py b/circe/cohortdefinition/builders/condition_era.py index 180480cb..02bedc8e 100644 --- a/circe/cohortdefinition/builders/condition_era.py +++ b/circe/cohortdefinition/builders/condition_era.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import ConditionEra from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -90,7 +88,7 @@ def embed_ordinal_expression(self, query: str, criteria: ConditionEra, where_cla def resolve_select_clauses( self, criteria: ConditionEra, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for condition era criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -120,7 +118,7 @@ def resolve_select_clauses( return select_cols def resolve_join_clauses( - self, criteria: ConditionEra, options: Optional[BuilderOptions] = None + self, criteria: ConditionEra, options: BuilderOptions | None = None ) -> list[str]: """Resolve join clauses for condition era criteria.""" join_clauses = [] @@ -139,7 +137,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: ConditionEra, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for condition era criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/condition_occurrence.py b/circe/cohortdefinition/builders/condition_occurrence.py index d9173fd7..e937f825 100644 --- a/circe/cohortdefinition/builders/condition_occurrence.py +++ b/circe/cohortdefinition/builders/condition_occurrence.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import ConditionOccurrence from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -104,7 +102,7 @@ def embed_ordinal_expression( def resolve_select_clauses( self, criteria: ConditionOccurrence, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for condition occurrence criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -158,7 +156,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: ConditionOccurrence, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for condition occurrence criteria.""" join_clauses = [] @@ -190,7 +188,7 @@ def resolve_join_clauses( return join_clauses def resolve_where_clauses( - self, criteria: ConditionOccurrence, options: Optional[BuilderOptions] = None + self, criteria: ConditionOccurrence, options: BuilderOptions | None = None ) -> list[str]: """Resolve where clauses for condition occurrence criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/death.py b/circe/cohortdefinition/builders/death.py index 2ae267cc..247ed352 100644 --- a/circe/cohortdefinition/builders/death.py +++ b/circe/cohortdefinition/builders/death.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import Death from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -81,7 +79,7 @@ def embed_ordinal_expression(self, query: str, criteria: Death, where_clauses: l """ return query - def resolve_select_clauses(self, criteria: Death, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_select_clauses(self, criteria: Death, options: BuilderOptions | None = None) -> list[str]: """Resolve select clauses for death criteria.""" select_cols = ["d.person_id", "d.cause_concept_id"] @@ -106,7 +104,7 @@ def resolve_select_clauses(self, criteria: Death, options: Optional[BuilderOptio return select_cols - def resolve_join_clauses(self, criteria: Death, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: Death, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for death criteria.""" joins = [] @@ -120,7 +118,7 @@ def resolve_join_clauses(self, criteria: Death, options: Optional[BuilderOptions return joins - def resolve_where_clauses(self, criteria: Death, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_where_clauses(self, criteria: Death, options: BuilderOptions | None = None) -> list[str]: """Resolve where clauses for death criteria.""" where_clauses = super().resolve_where_clauses(criteria) diff --git a/circe/cohortdefinition/builders/dose_era.py b/circe/cohortdefinition/builders/dose_era.py index f0b9c591..8835e1ae 100644 --- a/circe/cohortdefinition/builders/dose_era.py +++ b/circe/cohortdefinition/builders/dose_era.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import DoseEra from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -97,7 +95,7 @@ def embed_ordinal_expression(self, query: str, criteria: DoseEra, where_clauses: def resolve_select_clauses( self, criteria: DoseEra, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for dose era criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -124,7 +122,7 @@ def resolve_select_clauses( return select_cols - def resolve_join_clauses(self, criteria: DoseEra, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: DoseEra, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for dose era criteria.""" join_clauses = [] @@ -139,7 +137,7 @@ def resolve_join_clauses(self, criteria: DoseEra, options: Optional[BuilderOptio return join_clauses - def resolve_where_clauses(self, criteria: DoseEra, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_where_clauses(self, criteria: DoseEra, options: BuilderOptions | None = None) -> list[str]: """Resolve where clauses for dose era criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/drug_era.py b/circe/cohortdefinition/builders/drug_era.py index a44c9b4c..92aa7180 100644 --- a/circe/cohortdefinition/builders/drug_era.py +++ b/circe/cohortdefinition/builders/drug_era.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import DrugEra from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -101,7 +99,7 @@ def embed_ordinal_expression(self, query: str, criteria: DrugEra, where_clauses: def resolve_select_clauses( self, criteria: DrugEra, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for drug era criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -130,7 +128,7 @@ def resolve_select_clauses( return select_cols - def resolve_join_clauses(self, criteria: DrugEra, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: DrugEra, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for drug era criteria.""" join_clauses = [] @@ -145,7 +143,7 @@ def resolve_join_clauses(self, criteria: DrugEra, options: Optional[BuilderOptio return join_clauses - def resolve_where_clauses(self, criteria: DrugEra, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_where_clauses(self, criteria: DrugEra, options: BuilderOptions | None = None) -> list[str]: """Resolve where clauses for drug era criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/drug_exposure.py b/circe/cohortdefinition/builders/drug_exposure.py index 6ae1994c..188fdacc 100644 --- a/circe/cohortdefinition/builders/drug_exposure.py +++ b/circe/cohortdefinition/builders/drug_exposure.py @@ -9,8 +9,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import DrugExposure from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -122,7 +120,7 @@ def embed_ordinal_expression(self, query: str, criteria: DrugExposure, where_cla def resolve_select_clauses( self, criteria: DrugExposure, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for drug exposure criteria. @@ -192,7 +190,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: DrugExposure, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for drug exposure criteria. @@ -223,7 +221,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: DrugExposure, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for drug exposure criteria. diff --git a/circe/cohortdefinition/builders/location_region.py b/circe/cohortdefinition/builders/location_region.py index e305408c..d3aed822 100644 --- a/circe/cohortdefinition/builders/location_region.py +++ b/circe/cohortdefinition/builders/location_region.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import LocationRegion from .base import CriteriaSqlBuilder from .utils import BuilderOptions, CriteriaColumn @@ -80,7 +78,7 @@ def embed_ordinal_expression(self, query: str, criteria: LocationRegion, where_c def resolve_select_clauses( self, criteria: LocationRegion, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for location region criteria.""" # Default select columns that are always returned @@ -101,7 +99,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: LocationRegion, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for location region criteria.""" return [] @@ -109,7 +107,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: LocationRegion, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for location region criteria.""" return [] diff --git a/circe/cohortdefinition/builders/measurement.py b/circe/cohortdefinition/builders/measurement.py index f00f9b2c..3c66a17c 100644 --- a/circe/cohortdefinition/builders/measurement.py +++ b/circe/cohortdefinition/builders/measurement.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import Measurement from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -89,7 +87,7 @@ def embed_codeset_clause(self, query: str, criteria: Measurement) -> str: def resolve_select_clauses( self, criteria: Measurement, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for measurement criteria. @@ -155,7 +153,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: Measurement, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for measurement criteria. @@ -193,7 +191,7 @@ def resolve_join_clauses( def resolve_ordinal_expression( self, criteria: Measurement, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> str: """Resolve ordinal expression for measurement criteria.""" if criteria.first: @@ -203,7 +201,7 @@ def resolve_ordinal_expression( def resolve_where_clauses( self, criteria: Measurement, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for measurement criteria. diff --git a/circe/cohortdefinition/builders/observation.py b/circe/cohortdefinition/builders/observation.py index 8efdcb76..4c274e73 100644 --- a/circe/cohortdefinition/builders/observation.py +++ b/circe/cohortdefinition/builders/observation.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import Observation from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -72,7 +70,7 @@ def embed_codeset_clause(self, query: str, criteria: Observation) -> str: def resolve_select_clauses( self, criteria: Observation, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for observation criteria. @@ -110,9 +108,7 @@ def resolve_select_clauses( return select_cols - def resolve_join_clauses( - self, criteria: Observation, options: Optional[BuilderOptions] = None - ) -> list[str]: + def resolve_join_clauses(self, criteria: Observation, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for observation criteria. Java equivalent: ObservationSqlBuilder.resolveJoinClauses() @@ -149,7 +145,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: Observation, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for observation criteria.""" where_clauses = super().resolve_where_clauses(criteria) diff --git a/circe/cohortdefinition/builders/observation_period.py b/circe/cohortdefinition/builders/observation_period.py index 2ece7b56..c0854414 100644 --- a/circe/cohortdefinition/builders/observation_period.py +++ b/circe/cohortdefinition/builders/observation_period.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import ObservationPeriod from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -73,7 +71,7 @@ def get_table_column_for_criteria_column(self, criteria_column: CriteriaColumn) def get_criteria_sql_with_options( self, criteria: ObservationPeriod, - options: Optional[BuilderOptions], + options: BuilderOptions | None, ) -> str: """Get SQL query for criteria with builder options.""" query = super().get_criteria_sql_with_options(criteria, options) @@ -112,7 +110,7 @@ def embed_ordinal_expression( def resolve_select_clauses( self, criteria: ObservationPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for observation period criteria. @@ -148,7 +146,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: ObservationPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for observation period criteria.""" join_clauses = [] @@ -162,7 +160,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: ObservationPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for observation period criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/payer_plan_period.py b/circe/cohortdefinition/builders/payer_plan_period.py index 27466f75..1185c819 100644 --- a/circe/cohortdefinition/builders/payer_plan_period.py +++ b/circe/cohortdefinition/builders/payer_plan_period.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import PayerPlanPeriod from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -77,7 +75,7 @@ def get_table_column_for_criteria_column(self, criteria_column: CriteriaColumn) def get_criteria_sql_with_options( self, criteria: PayerPlanPeriod, - options: Optional[BuilderOptions], + options: BuilderOptions | None, ) -> str: """Get SQL query for criteria with builder options.""" query = super().get_criteria_sql_with_options(criteria, options) @@ -115,7 +113,7 @@ def embed_ordinal_expression( def resolve_select_clauses( self, criteria: PayerPlanPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for payer plan period criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -184,7 +182,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: PayerPlanPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for payer plan period criteria.""" join_clauses = [] @@ -202,7 +200,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: PayerPlanPeriod, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for payer plan period criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/procedure_occurrence.py b/circe/cohortdefinition/builders/procedure_occurrence.py index de1d7bc4..864836f4 100644 --- a/circe/cohortdefinition/builders/procedure_occurrence.py +++ b/circe/cohortdefinition/builders/procedure_occurrence.py @@ -9,8 +9,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import Criteria from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -129,7 +127,7 @@ def embed_codeset_clause(self, query: str, criteria: Criteria) -> str: def resolve_select_clauses( self, criteria: Criteria, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for criteria. @@ -183,7 +181,7 @@ def resolve_select_clauses( return select_cols - def resolve_join_clauses(self, criteria: Criteria, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: Criteria, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for criteria. Java equivalent: ProcedureOccurrenceSqlBuilder.resolveJoinClauses() @@ -221,7 +219,7 @@ def resolve_join_clauses(self, criteria: Criteria, options: Optional[BuilderOpti def resolve_where_clauses( self, criteria: Criteria, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for criteria. diff --git a/circe/cohortdefinition/builders/specimen.py b/circe/cohortdefinition/builders/specimen.py index 7aaea3d5..a09fbcac 100644 --- a/circe/cohortdefinition/builders/specimen.py +++ b/circe/cohortdefinition/builders/specimen.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import Specimen from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -82,7 +80,7 @@ def embed_ordinal_expression(self, query: str, criteria: Specimen, where_clauses query = query.replace("@ordinalExpression", "") return query - def resolve_join_clauses(self, criteria: Specimen, options: Optional[BuilderOptions] = None) -> list[str]: + def resolve_join_clauses(self, criteria: Specimen, options: BuilderOptions | None = None) -> list[str]: """Resolve join clauses for specimen criteria.""" joins = [] @@ -99,7 +97,7 @@ def resolve_join_clauses(self, criteria: Specimen, options: Optional[BuilderOpti def resolve_where_clauses( self, criteria: Specimen, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for specimen criteria.""" where_clauses = [] diff --git a/circe/cohortdefinition/builders/utils.py b/circe/cohortdefinition/builders/utils.py index 5ef033fd..03af8506 100644 --- a/circe/cohortdefinition/builders/utils.py +++ b/circe/cohortdefinition/builders/utils.py @@ -9,7 +9,7 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Any, Optional +from typing import Any from ...vocabulary.concept import Concept from ..core import DateAdjustment, DateRange, NumericRange @@ -62,9 +62,9 @@ def get_date_adjustment_expression( @staticmethod def get_codeset_join_expression( - standard_codeset_id: Optional[int], + standard_codeset_id: int | None, standard_concept_column: str, - source_codeset_id: Optional[int], + source_codeset_id: int | None, source_concept_column: str, ) -> str: """Get codeset join expression for SQL. @@ -133,7 +133,7 @@ def get_operator(op: str) -> str: raise RuntimeError(f"Unknown operator type: {op}") @staticmethod - def build_date_range_clause(sql_expression: str, date_range: Optional[DateRange]) -> Optional[str]: + def build_date_range_clause(sql_expression: str, date_range: DateRange | None) -> str | None: """Build date range clause for SQL. Java equivalent: BuilderUtils.buildDateRangeClause(String sqlExpression, DateRange range) @@ -157,9 +157,9 @@ def build_date_range_clause(sql_expression: str, date_range: Optional[DateRange] @staticmethod def build_numeric_range_clause( sql_expression: str, - numeric_range: Optional[NumericRange], - format: Optional[str] = None, - ) -> Optional[str]: + numeric_range: NumericRange | None, + format: str | None = None, + ) -> str | None: """Build numeric range clause for SQL. Java equivalent: BuilderUtils.buildNumericRangeClause(String sqlExpression, NumericRange range, String format) @@ -193,7 +193,7 @@ def build_numeric_range_clause( return f"{sql_expression} {BuilderUtils.get_operator(op)} {int(numeric_range.value)}" @staticmethod - def build_text_filter_clause(text_filter: Optional[Any], column_name: str) -> Optional[str]: + def build_text_filter_clause(text_filter: Any | None, column_name: str) -> str | None: """Build text filter clause for SQL. Java equivalent: BuilderUtils.buildTextFilterClause() diff --git a/circe/cohortdefinition/builders/visit_detail.py b/circe/cohortdefinition/builders/visit_detail.py index 8eb93569..813b8050 100644 --- a/circe/cohortdefinition/builders/visit_detail.py +++ b/circe/cohortdefinition/builders/visit_detail.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import VisitDetail from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -102,7 +100,7 @@ def embed_ordinal_expression(self, query: str, criteria: VisitDetail, where_clau def resolve_select_clauses( self, criteria: VisitDetail, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for visit detail criteria.""" select_cols = list(self.DEFAULT_SELECT_COLUMNS) @@ -151,7 +149,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: VisitDetail, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for visit detail criteria.""" join_clauses = [] @@ -177,7 +175,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: VisitDetail, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for visit detail criteria.""" where_clauses = [] @@ -264,7 +262,7 @@ def add_where_clause( where_clauses: list[str], concept_set_selection, concept_column: str, - exclude: Optional[bool] = None, + exclude: bool | None = None, ): """Add where clause for concept set selection.""" is_exclusion = exclude if exclude is not None else concept_set_selection.is_exclusion diff --git a/circe/cohortdefinition/builders/visit_occurrence.py b/circe/cohortdefinition/builders/visit_occurrence.py index b68ec3ff..1620035c 100644 --- a/circe/cohortdefinition/builders/visit_occurrence.py +++ b/circe/cohortdefinition/builders/visit_occurrence.py @@ -8,8 +8,6 @@ Reference: JAVA_CLASS_MAPPINGS.md for Java equivalents. """ -from typing import Optional - from ..criteria import VisitOccurrence from .base import CriteriaSqlBuilder from .utils import BuilderOptions, BuilderUtils, CriteriaColumn @@ -74,7 +72,7 @@ def embed_codeset_clause(self, query: str, criteria: VisitOccurrence) -> str: def resolve_select_clauses( self, criteria: VisitOccurrence, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve select clauses for visit occurrence criteria.""" # Default select columns that are always returned @@ -125,7 +123,7 @@ def resolve_select_clauses( def resolve_join_clauses( self, criteria: VisitOccurrence, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve join clauses for visit occurrence criteria.""" join_clauses = [] @@ -162,7 +160,7 @@ def resolve_join_clauses( def resolve_where_clauses( self, criteria: VisitOccurrence, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> list[str]: """Resolve where clauses for visit occurrence criteria.""" where_clauses = super().resolve_where_clauses(criteria, options) diff --git a/circe/cohortdefinition/cohort.py b/circe/cohortdefinition/cohort.py index 6071dc81..b290d83d 100644 --- a/circe/cohortdefinition/cohort.py +++ b/circe/cohortdefinition/cohort.py @@ -9,8 +9,7 @@ """ import contextlib -import json -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any from pydantic import ( AliasChoices, @@ -51,6 +50,32 @@ InclusionRule = Any +def _python_serialize(obj: Any) -> bytes: + """Deterministic Python-native serialization for checksum hashing. + + Recursively serializes Python builtins (dict, list, str, int, float, + bool, None) to a stable byte representation. Dict keys are sorted to + guarantee deterministic output across Python versions and platforms. + """ + if isinstance(obj, dict): + items = b",".join(_python_serialize(k) + b":" + _python_serialize(v) for k, v in sorted(obj.items())) + return b"{" + items + b"}" + if isinstance(obj, list): + items = b",".join(_python_serialize(v) for v in obj) + return b"[" + items + b"]" + if isinstance(obj, bool): + return b"true" if obj else b"false" + if isinstance(obj, int): + return repr(obj).encode("ascii") + if isinstance(obj, float): + return repr(obj).encode("ascii") + if isinstance(obj, str): + return obj.encode("utf-8") + if obj is None: + return b"null" + return repr(obj).encode("utf-8") + + class CohortExpression(CirceBaseModel): """Main cohort expression class containing all cohort definition components. @@ -62,38 +87,38 @@ class CohortExpression(CirceBaseModel): validation_alias=AliasChoices("ConceptSets", "conceptSets"), serialization_alias="ConceptSets", ) - qualified_limit: Optional[ResultLimit] = Field( + qualified_limit: ResultLimit | None = Field( default=None, validation_alias=AliasChoices("QualifiedLimit", "qualifiedLimit"), serialization_alias="QualifiedLimit", ) - additional_criteria: Optional[CriteriaGroup] = Field( + additional_criteria: CriteriaGroup | None = Field( default=None, validation_alias=AliasChoices("AdditionalCriteria", "additionalCriteria"), serialization_alias="AdditionalCriteria", ) - end_strategy: Optional[Union[EndStrategy, DateOffsetStrategy, CustomEraStrategy]] = Field( + end_strategy: EndStrategy | DateOffsetStrategy | CustomEraStrategy | None = Field( default=None, validation_alias=AliasChoices("EndStrategy", "endStrategy"), serialization_alias="EndStrategy", ) - cdm_version_range: Optional[str] = Field(default=None, alias="cdmVersionRange") - primary_criteria: Optional[PrimaryCriteria] = Field( + cdm_version_range: str | None = Field(default=None, alias="cdmVersionRange") + primary_criteria: PrimaryCriteria | None = Field( default=None, validation_alias=AliasChoices("PrimaryCriteria", "primaryCriteria"), serialization_alias="PrimaryCriteria", ) - expression_limit: Optional[ResultLimit] = Field( + expression_limit: ResultLimit | None = Field( default=None, validation_alias=AliasChoices("ExpressionLimit", "expressionLimit"), serialization_alias="ExpressionLimit", ) - collapse_settings: Optional[CollapseSettings] = Field( + collapse_settings: CollapseSettings | None = Field( default=None, validation_alias=AliasChoices("CollapseSettings", "collapseSettings"), serialization_alias="CollapseSettings", ) - title: Optional[str] = Field( + title: str | None = Field( default=None, validation_alias=AliasChoices("Title", "title"), serialization_alias="Title", @@ -103,7 +128,7 @@ class CohortExpression(CirceBaseModel): validation_alias=AliasChoices("InclusionRules", "inclusionRules"), serialization_alias="InclusionRules", ) - censor_window: Optional[Period] = Field( + censor_window: Period | None = Field( default=None, validation_alias=AliasChoices("CensorWindow", "censorWindow"), serialization_alias="CensorWindow", @@ -364,19 +389,12 @@ def checksum(self, algorithm: str = "sha256") -> str: Hex digest of the checksum """ import hashlib - import json - - # 1. Dump with defaults excluded to handle implicit defaults - data = self.model_dump(exclude_unset=True, exclude_defaults=True, by_alias=True) - - # 2. Normalize: remove metadata, deduplicate concept sets, etc. - normalized_data = self._normalize_for_checksum(data) - - # 3. Serialize to canonical JSON - canonical_json = json.dumps(normalized_data, sort_keys=True) + data = self.model_dump(by_alias=True, exclude_none=True) + normalized = self._normalize_for_checksum(data) + serialized = _python_serialize(normalized) h = hashlib.new(algorithm) - h.update(canonical_json.encode("utf-8")) + h.update(serialized) return h.hexdigest() def _normalize_for_checksum(self, data: Any) -> Any: @@ -396,20 +414,14 @@ def _normalize_for_checksum(self, data: Any) -> Any: seen_items = set() for item in data["items"]: - # Normalize the item first norm_item = self._normalize_for_checksum(item) + item_bytes = _python_serialize(norm_item) - # Create a sortable/hashable representation for deduplication - # We need to sort keys to ensure tuple order is consistent - item_json = json.dumps(norm_item, sort_keys=True) - - if item_json not in seen_items: - seen_items.add(item_json) + if item_bytes not in seen_items: + seen_items.add(item_bytes) normalized_items.append(norm_item) - # Sort items to ensure list order doesn't affect hash - # Sort by the JSON string representation - normalized_items.sort(key=lambda x: json.dumps(x, sort_keys=True)) + normalized_items.sort(key=_python_serialize) new_data = data.copy() new_data["items"] = normalized_items @@ -417,8 +429,6 @@ def _normalize_for_checksum(self, data: Any) -> Any: # Handle Concept Objects (heuristically by fields) if "CONCEPT_ID" in data: - # Keep ID, remove metadata names/codes/vocab - # Keep only structural identifier return {"CONCEPT_ID": data["CONCEPT_ID"]} # Recurse for other dicts @@ -524,7 +534,7 @@ def has_end_strategy(self) -> bool: """ return self.end_strategy is not None - def get_end_strategy_type(self) -> Optional[str]: + def get_end_strategy_type(self) -> str | None: """Get the type of end strategy. Returns: @@ -563,7 +573,7 @@ def has_observation_window(self) -> bool: return self.primary_criteria.observation_window is not None - def get_primary_limit_type(self) -> Optional[str]: + def get_primary_limit_type(self) -> str | None: """Get the primary limit type. Returns: diff --git a/circe/cohortdefinition/cohort_expression_query_builder.py b/circe/cohortdefinition/cohort_expression_query_builder.py index b275ce6c..254bff0d 100644 --- a/circe/cohortdefinition/cohort_expression_query_builder.py +++ b/circe/cohortdefinition/cohort_expression_query_builder.py @@ -9,7 +9,7 @@ """ import json -from typing import Any, Optional, Union +from typing import Any from circe.extensions import get_registry @@ -69,12 +69,12 @@ class BuildExpressionQueryOptions: """ def __init__(self): - self.cohort_id_field_name: Optional[str] = None - self.cohort_id: Optional[int] = None - self.cdm_schema: Optional[str] = None - self.target_table: Optional[str] = None - self.result_schema: Optional[str] = None - self.vocabulary_schema: Optional[str] = None + self.cohort_id_field_name: str | None = None + self.cohort_id: int | None = None + self.cdm_schema: str | None = None + self.target_table: str | None = None + self.result_schema: str | None = None + self.vocabulary_schema: str | None = None self.generate_stats: bool = False @classmethod @@ -592,7 +592,7 @@ def get_censoring_events_query(self, censoring_criteria: list[Criteria]) -> str: def get_primary_events_query( self, primary_criteria: PrimaryCriteria, - subquery: Optional[str] = None, + subquery: str | None = None, ) -> str: """Get primary events query. @@ -652,7 +652,7 @@ def _get_primary_events_subquery(self, primary_criteria: PrimaryCriteria) -> str return query - def get_final_cohort_query(self, censor_window: Optional[Period]) -> str: + def get_final_cohort_query(self, censor_window: Period | None) -> str: """Get final cohort query. Java equivalent: getFinalCohortQuery() @@ -780,7 +780,7 @@ def _build_inclusion_analysis_section(self, expression: CohortExpression) -> str def build_expression_query( self, - expression: Union[str, CohortExpression], + expression: str | CohortExpression, options: BuildExpressionQueryOptions, ) -> str: """Build expression query from CohortExpression object or JSON string. @@ -1170,7 +1170,7 @@ def _get_windowed_criteria_query_internal( sql_template: str, criteria: Any, event_table: str, - options: Optional[BuilderOptions], + options: BuilderOptions | None, ) -> str: """Get windowed criteria query (internal method with all parameters). @@ -1365,7 +1365,7 @@ def get_windowed_criteria_query( self, criteria: Any, event_table: str, - options: Optional[BuilderOptions] = None, + options: BuilderOptions | None = None, ) -> str: """Get windowed criteria query. @@ -1465,7 +1465,7 @@ def get_corelated_criteria_query(self, corelated_criteria: CorelatedCriteria, ev return query - def get_criteria_sql(self, criteria: Criteria, options: Optional[BuilderOptions] = None) -> str: + def get_criteria_sql(self, criteria: Criteria, options: BuilderOptions | None = None) -> str: """Get criteria SQL for any criteria type. Java equivalent: Various getCriteriaSql methods @@ -1614,7 +1614,7 @@ def _get_criteria_sql_from_builder( self, builder: Any, criteria: Criteria, - options: Optional[BuilderOptions], + options: BuilderOptions | None, ) -> str: """Generic method to get criteria SQL from builder.""" query = builder.get_criteria_sql_with_options(criteria, options) @@ -1637,7 +1637,7 @@ def get_date_field_for_offset_strategy(self, date_field: str) -> str: def get_strategy_sql( self, - strategy: Union[DateOffsetStrategy, CustomEraStrategy], + strategy: DateOffsetStrategy | CustomEraStrategy, event_table: str, ) -> str: """Get strategy SQL for date offset or custom era strategy.""" diff --git a/circe/cohortdefinition/core.py b/circe/cohortdefinition/core.py index 714da872..b5dd7ec1 100644 --- a/circe/cohortdefinition/core.py +++ b/circe/cohortdefinition/core.py @@ -9,7 +9,7 @@ """ from enum import Enum -from typing import Any, Optional, Union +from typing import Any from pydantic import ( AliasChoices, @@ -91,7 +91,7 @@ class ResultLimit(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.ResultLimit """ - type: Optional[str] = Field( + type: str | None = Field( default=None, validation_alias=AliasChoices("Type", "type"), serialization_alias="Type", @@ -104,8 +104,8 @@ class Period(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.Period """ - start_date: Optional[str] = None - end_date: Optional[str] = None + start_date: str | None = None + end_date: str | None = None model_config = ConfigDict(populate_by_name=True, alias_generator=to_pascal_alias) @@ -116,17 +116,17 @@ class DateRange(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.DateRange """ - op: Optional[str] = Field( + op: str | None = Field( default=None, validation_alias=AliasChoices("Op", "op"), serialization_alias="Op", ) - value: Optional[Union[str, float]] = Field( + value: str | float | None = Field( default=None, validation_alias=AliasChoices("Value", "value"), serialization_alias="Value", ) - extent: Optional[Union[str, float]] = Field( + extent: str | float | None = Field( default=None, validation_alias=AliasChoices("Extent", "extent"), serialization_alias="Extent", @@ -139,17 +139,17 @@ class NumericRange(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.NumericRange """ - op: Optional[str] = Field( + op: str | None = Field( default=None, validation_alias=AliasChoices("Op", "op"), serialization_alias="Op", ) - value: Optional[Union[int, float]] = Field( + value: int | float | None = Field( default=None, validation_alias=AliasChoices("Value", "value"), serialization_alias="Value", ) - extent: Optional[Union[int, float]] = Field( + extent: int | float | None = Field( default=None, validation_alias=AliasChoices("Extent", "extent"), serialization_alias="Extent", @@ -170,12 +170,12 @@ class DateAdjustment(CirceBaseModel): validation_alias=AliasChoices("endOffset", "EndOffset"), serialization_alias="endOffset", ) - start_with: Optional[DateType] = Field( + start_with: DateType | None = Field( default=DateType.START_DATE, validation_alias=AliasChoices("startWith", "StartWith"), serialization_alias="startWith", ) - end_with: Optional[DateType] = Field( + end_with: DateType | None = Field( default=DateType.END_DATE, validation_alias=AliasChoices("endWith", "EndWith"), serialization_alias="endWith", @@ -209,7 +209,7 @@ class CollapseSettings(CirceBaseModel): """ era_pad: int = Field(validation_alias=AliasChoices("EraPad", "eraPad"), serialization_alias="EraPad") - collapse_type: Optional[CollapseType] = Field( + collapse_type: CollapseType | None = Field( default=CollapseType.ERA, validation_alias=AliasChoices("CollapseType", "collapseType"), serialization_alias="CollapseType", @@ -224,7 +224,7 @@ class EndStrategy(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.EndStrategy """ - include: Optional[str] = None # JsonTypeInfo.Id.NAME + include: str | None = None # JsonTypeInfo.Id.NAME @model_serializer(mode="wrap") def _serialize_polymorphic(self, serializer, info): @@ -243,7 +243,7 @@ class ConceptSetSelection(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.ConceptSetSelection """ - codeset_id: Optional[int] = Field( + codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("CodesetId", "codesetId"), serialization_alias="CodesetId", @@ -263,12 +263,12 @@ class TextFilter(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.TextFilter """ - text: Optional[str] = Field( + text: str | None = Field( default=None, validation_alias=AliasChoices("Text", "text"), serialization_alias="Text", ) - op: Optional[str] = Field( + op: str | None = Field( default=None, validation_alias=AliasChoices("Op", "op"), serialization_alias="Op", @@ -282,7 +282,7 @@ class WindowBound(CirceBaseModel): """ coeff: int = Field(validation_alias=AliasChoices("Coeff", "coeff"), serialization_alias="Coeff") - days: Optional[int] = Field( + days: int | None = Field( default=None, validation_alias=AliasChoices("Days", "days"), serialization_alias="Days", @@ -297,22 +297,22 @@ class Window(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.Window """ - start: Optional[WindowBound] = Field( + start: WindowBound | None = Field( default=None, validation_alias=AliasChoices("Start", "start"), serialization_alias="Start", ) - end: Optional[WindowBound] = Field( + end: WindowBound | None = Field( default=None, validation_alias=AliasChoices("End", "end"), serialization_alias="End", ) - use_event_end: Optional[bool] = Field( + use_event_end: bool | None = Field( default=None, validation_alias=AliasChoices("UseEventEnd", "useEventEnd"), serialization_alias="UseEventEnd", ) - use_index_end: Optional[bool] = Field( + use_index_end: bool | None = Field( default=None, validation_alias=AliasChoices("UseIndexEnd", "useIndexEnd"), serialization_alias="UseIndexEnd", @@ -346,7 +346,7 @@ class CustomEraStrategy(EndStrategy): Java equivalent: org.ohdsi.circe.cohortdefinition.CustomEraStrategy """ - drug_codeset_id: Optional[int] = Field( + drug_codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("DrugCodesetId", "drugCodesetId"), serialization_alias="DrugCodesetId", @@ -361,7 +361,7 @@ class CustomEraStrategy(EndStrategy): validation_alias=AliasChoices("Offset", "offset"), serialization_alias="Offset", ) - days_supply_override: Optional[int] = Field( + days_supply_override: int | None = Field( default=None, validation_alias=AliasChoices("DaysSupplyOverride", "daysSupplyOverride"), serialization_alias="DaysSupplyOverride", diff --git a/circe/cohortdefinition/criteria.py b/circe/cohortdefinition/criteria.py index d1542b10..692b0914 100644 --- a/circe/cohortdefinition/criteria.py +++ b/circe/cohortdefinition/criteria.py @@ -9,7 +9,7 @@ """ from enum import Enum -from typing import Annotated, Any, Optional, Union +from typing import Annotated, Any, Optional from pydantic import ( AliasChoices, @@ -109,7 +109,7 @@ class Occurrence(CirceBaseModel): validation_alias=AliasChoices("IsDistinct", "isDistinct"), serialization_alias="IsDistinct", ) - count_column: Optional[CriteriaColumn] = Field( + count_column: CriteriaColumn | None = Field( default=None, validation_alias=AliasChoices("CountColumn", "countColumn"), serialization_alias="CountColumn", @@ -135,12 +135,12 @@ class WindowedCriteria(CirceBaseModel): validation_alias=AliasChoices("Criteria", "criteria"), serialization_alias="Criteria", ) - start_window: Optional[Window] = Field( + start_window: Window | None = Field( default=None, validation_alias=AliasChoices("StartWindow", "startWindow"), serialization_alias="StartWindow", ) - end_window: Optional[Window] = Field( + end_window: Window | None = Field( default=None, validation_alias=AliasChoices("EndWindow", "endWindow"), serialization_alias="EndWindow", @@ -166,7 +166,7 @@ class CorelatedCriteria(WindowedCriteria): Java equivalent: org.ohdsi.circe.cohortdefinition.CorelatedCriteria """ - occurrence: Optional[Occurrence] = Field( + occurrence: Occurrence | None = Field( default=None, validation_alias=AliasChoices("Occurrence", "occurrence"), serialization_alias="Occurrence", @@ -180,47 +180,47 @@ class DemographicCriteria(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.DemographicCriteria """ - gender: Optional[list[Concept]] = Field( + gender: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("Gender", "gender"), serialization_alias="Gender", ) - occurrence_end_date: Optional[DateRange] = Field( + occurrence_end_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceEndDate", "occurrenceEndDate"), serialization_alias="OccurrenceEndDate", ) - gender_cs: Optional[ConceptSetSelection] = Field( + gender_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("GenderCS", "genderCS"), serialization_alias="GenderCS", ) - race: Optional[list[Concept]] = Field( + race: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("Race", "race"), serialization_alias="Race", ) - ethnicity_cs: Optional[ConceptSetSelection] = Field( + ethnicity_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("EthnicityCS", "ethnicityCS"), serialization_alias="EthnicityCS", ) - age: Optional[NumericRange] = Field( + age: NumericRange | None = Field( default=None, validation_alias=AliasChoices("Age", "age"), serialization_alias="Age", ) - race_cs: Optional[ConceptSetSelection] = Field( + race_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("RaceCS", "raceCS"), serialization_alias="RaceCS", ) - ethnicity: Optional[list[Concept]] = Field( + ethnicity: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("Ethnicity", "ethnicity"), serialization_alias="Ethnicity", ) - occurrence_start_date: Optional[DateRange] = Field( + occurrence_start_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceStartDate", "occurrenceStartDate"), serialization_alias="OccurrenceStartDate", @@ -235,7 +235,7 @@ class Criteria(CirceBaseModel): Java equivalent: org.ohdsi.circe.cohortdefinition.Criteria """ - date_adjustment: Optional[DateAdjustment] = Field( + date_adjustment: DateAdjustment | None = Field( default=None, validation_alias=AliasChoices("DateAdjustment", "dateAdjustment"), serialization_alias="DateAdjustment", @@ -245,7 +245,7 @@ class Criteria(CirceBaseModel): validation_alias=AliasChoices("CorrelatedCriteria", "correlatedCriteria"), serialization_alias="CorrelatedCriteria", ) - include: Optional[str] = None # JsonTypeInfo.Id.NAME + include: str | None = None # JsonTypeInfo.Id.NAME @model_serializer(mode="wrap") def _serialize_polymorphic(self, serializer, info): @@ -271,7 +271,7 @@ def _serialize_polymorphic(self, serializer, info): return {self.__class__.__name__: data} - def accept(self, dispatcher: Any, options: Optional[Any] = None) -> str: + def accept(self, dispatcher: Any, options: Any | None = None) -> str: """Accept method for visitor pattern.""" return dispatcher.get_criteria_sql(self, options) @@ -287,12 +287,12 @@ class InclusionRule(CirceBaseModel): validation_alias=AliasChoices("Expression", "expression"), serialization_alias="Expression", ) - description: Optional[str] = Field( + description: str | None = Field( default=None, validation_alias=AliasChoices("Description", "description"), serialization_alias="Description", ) - name: Optional[str] = Field( + name: str | None = Field( default=None, validation_alias=AliasChoices("Name", "name"), serialization_alias="Name", @@ -310,93 +310,93 @@ class ConditionOccurrence(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.ConditionOccurrence """ - codeset_id: Optional[int] = Field( + codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("CodesetId", "codesetId"), serialization_alias="CodesetId", ) - first: Optional[bool] = Field( + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - occurrence_start_date: Optional[DateRange] = Field( + occurrence_start_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceStartDate", "occurrenceStartDate"), serialization_alias="OccurrenceStartDate", ) - occurrence_end_date: Optional[DateRange] = Field( + occurrence_end_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceEndDate", "occurrenceEndDate"), serialization_alias="OccurrenceEndDate", ) - condition_type: Optional[list[Concept]] = Field( + condition_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ConditionType", "conditionType"), serialization_alias="ConditionType", ) - condition_type_cs: Optional[ConceptSetSelection] = Field( + condition_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ConditionTypeCS", "conditionTypeCS"), serialization_alias="ConditionTypeCS", ) - condition_type_exclude: Optional[bool] = Field( + condition_type_exclude: bool | None = Field( default=False, validation_alias=AliasChoices("ConditionTypeExclude", "conditionTypeExclude"), serialization_alias="ConditionTypeExclude", ) - stop_reason: Optional[TextFilter] = Field( + stop_reason: TextFilter | None = Field( default=None, validation_alias=AliasChoices("StopReason", "stopReason"), serialization_alias="StopReason", ) - condition_source_concept: Optional[int] = Field( + condition_source_concept: int | None = Field( default=None, validation_alias=AliasChoices("ConditionSourceConcept", "conditionSourceConcept"), serialization_alias="ConditionSourceConcept", ) - age: Optional[NumericRange] = Field( + age: NumericRange | None = Field( default=None, validation_alias=AliasChoices("Age", "age"), serialization_alias="Age", ) - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field( + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("GenderCS", "genderCS"), serialization_alias="GenderCS", ) - provider_specialty: Optional[list[Concept]] = Field( + provider_specialty: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialty", "providerSpecialty"), serialization_alias="ProviderSpecialty", ) - provider_specialty_cs: Optional[ConceptSetSelection] = Field( + provider_specialty_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialtyCS", "providerSpecialtyCS"), serialization_alias="ProviderSpecialtyCS", ) - visit_type: Optional[list[Concept]] = Field( + visit_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("VisitType", "visitType"), serialization_alias="VisitType", ) - visit_type_cs: Optional[ConceptSetSelection] = Field( + visit_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("VisitTypeCS", "visitTypeCS"), serialization_alias="VisitTypeCS", ) - condition_status: Optional[list[Concept]] = Field( + condition_status: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ConditionStatus", "conditionStatus"), serialization_alias="ConditionStatus", ) - condition_status_cs: Optional[ConceptSetSelection] = Field( + condition_status_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ConditionStatusCS", "conditionStatusCS"), serialization_alias="ConditionStatusCS", ) - date_adjustment: Optional[DateAdjustment] = Field( + date_adjustment: DateAdjustment | None = Field( default=None, validation_alias=AliasChoices("DateAdjustment", "dateAdjustment"), serialization_alias="DateAdjustment", @@ -411,33 +411,33 @@ class DrugExposure(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.DrugExposure """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field( + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceEndDate", "occurrenceEndDate"), serialization_alias="OccurrenceEndDate", ) - stop_reason: Optional[TextFilter] = Field( + stop_reason: TextFilter | None = Field( default=None, validation_alias=AliasChoices("StopReason", "stopReason"), serialization_alias="StopReason", ) - drug_source_concept: Optional[int] = Field( + drug_source_concept: int | None = Field( default=None, validation_alias=AliasChoices("DrugSourceConcept", "drugSourceConcept"), serialization_alias="DrugSourceConcept", ) - gender_cs: Optional[ConceptSetSelection] = Field( + gender_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("GenderCS", "genderCS"), serialization_alias="GenderCS", ) - drug_type: Optional[list[Concept]] = Field( + drug_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("DrugType", "drugType"), serialization_alias="DrugType", ) - drug_type_cs: Optional[ConceptSetSelection] = Field( + drug_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("DrugTypeCS", "drugTypeCS"), serialization_alias="DrugTypeCS", @@ -447,83 +447,83 @@ class DrugExposure(Criteria): validation_alias=AliasChoices("DrugTypeExclude", "drugTypeExclude"), serialization_alias="DrugTypeExclude", ) - provider_specialty_cs: Optional[ConceptSetSelection] = Field( + provider_specialty_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialtyCS", "providerSpecialtyCS"), serialization_alias="ProviderSpecialtyCS", ) - visit_type_cs: Optional[ConceptSetSelection] = Field( + visit_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("VisitTypeCS", "visitTypeCS"), serialization_alias="VisitTypeCS", ) - visit_type: Optional[list[Concept]] = Field( + visit_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("VisitType", "visitType"), serialization_alias="VisitType", ) - route_concept: Optional[list[Concept]] = Field( + route_concept: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("RouteConcept", "routeConcept"), serialization_alias="RouteConcept", ) - route_concept_cs: Optional[ConceptSetSelection] = Field( + route_concept_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("RouteConceptCS", "routeConceptCS"), serialization_alias="RouteConceptCS", ) - codeset_id: Optional[int] = Field( + codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("CodesetId", "codesetId"), serialization_alias="CodesetId", ) - first: Optional[bool] = Field( + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - provider_specialty: Optional[list[Concept]] = Field( + provider_specialty: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialty", "providerSpecialty"), serialization_alias="ProviderSpecialty", ) - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field( + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceStartDate", "occurrenceStartDate"), serialization_alias="OccurrenceStartDate", ) - dose_unit: Optional[list[Concept]] = Field( + dose_unit: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("DoseUnit", "doseUnit"), serialization_alias="DoseUnit", ) - dose_unit_cs: Optional[ConceptSetSelection] = Field( + dose_unit_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("DoseUnitCS", "doseUnitCS"), serialization_alias="DoseUnitCS", ) - lot_number: Optional[TextFilter] = Field( + lot_number: TextFilter | None = Field( default=None, validation_alias=AliasChoices("LotNumber", "lotNumber"), serialization_alias="LotNumber", ) - quantity: Optional[NumericRange] = Field( + quantity: NumericRange | None = Field( default=None, validation_alias=AliasChoices("Quantity", "quantity"), serialization_alias="Quantity", ) - days_supply: Optional[NumericRange] = Field( + days_supply: NumericRange | None = Field( default=None, validation_alias=AliasChoices("DaysSupply", "daysSupply"), serialization_alias="DaysSupply", ) - refills: Optional[NumericRange] = Field( + refills: NumericRange | None = Field( default=None, validation_alias=AliasChoices("Refills", "refills"), serialization_alias="Refills", ) - effective_drug_dose: Optional[NumericRange] = Field( + effective_drug_dose: NumericRange | None = Field( default=None, validation_alias=AliasChoices("EffectiveDrugDose", "effectiveDrugDose"), serialization_alias="EffectiveDrugDose", @@ -538,28 +538,28 @@ class ProcedureOccurrence(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.ProcedureOccurrence """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - procedure_source_concept: Optional[int] = Field(default=None, alias="ProcedureSourceConcept") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - procedure_type: Optional[list[Concept]] = Field(default=None, alias="ProcedureType") - procedure_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProcedureTypeCS") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + procedure_source_concept: int | None = Field(default=None, alias="ProcedureSourceConcept") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + procedure_type: list[Concept] | None = Field(default=None, alias="ProcedureType") + procedure_type_cs: ConceptSetSelection | None = Field(default=None, alias="ProcedureTypeCS") procedure_type_exclude: bool = Field(default=False, alias="ProcedureTypeExclude") - provider_specialty_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProviderSpecialtyCS") - visit_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="VisitTypeCS") - visit_type: Optional[list[Concept]] = Field(default=None, alias="VisitType") - modifier: Optional[list[Concept]] = Field(default=None, alias="Modifier") - modifier_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ModifierCS") - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field( + provider_specialty_cs: ConceptSetSelection | None = Field(default=None, alias="ProviderSpecialtyCS") + visit_type_cs: ConceptSetSelection | None = Field(default=None, alias="VisitTypeCS") + visit_type: list[Concept] | None = Field(default=None, alias="VisitType") + modifier: list[Concept] | None = Field(default=None, alias="Modifier") + modifier_cs: ConceptSetSelection | None = Field(default=None, alias="ModifierCS") + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - age: Optional[NumericRange] = None - quantity: Optional[NumericRange] = Field(default=None, alias="Quantity") - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + age: NumericRange | None = None + quantity: NumericRange | None = Field(default=None, alias="Quantity") + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -570,23 +570,23 @@ class VisitOccurrence(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.VisitOccurrence """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field(default=None, alias="First") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - visit_type: Optional[list[Concept]] = Field(default=None, alias="VisitType") - visit_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="VisitTypeCS") + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field(default=None, alias="First") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + visit_type: list[Concept] | None = Field(default=None, alias="VisitType") + visit_type_cs: ConceptSetSelection | None = Field(default=None, alias="VisitTypeCS") visit_type_exclude: bool = Field(default=False, alias="VisitTypeExclude") - visit_source_concept: Optional[int] = Field(default=None, alias="VisitSourceConcept") - visit_length: Optional[NumericRange] = Field(default=None, alias="VisitLength") - provider_specialty_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProviderSpecialtyCS") - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - place_of_service: Optional[list[Concept]] = Field(default=None, alias="PlaceOfService") - place_of_service_cs: Optional[ConceptSetSelection] = Field(default=None, alias="PlaceOfServiceCS") - place_of_service_location: Optional[int] = Field(default=None, alias="PlaceOfServiceLocation") - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + visit_source_concept: int | None = Field(default=None, alias="VisitSourceConcept") + visit_length: NumericRange | None = Field(default=None, alias="VisitLength") + provider_specialty_cs: ConceptSetSelection | None = Field(default=None, alias="ProviderSpecialtyCS") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + place_of_service: list[Concept] | None = Field(default=None, alias="PlaceOfService") + place_of_service_cs: ConceptSetSelection | None = Field(default=None, alias="PlaceOfServiceCS") + place_of_service_location: int | None = Field(default=None, alias="PlaceOfServiceLocation") + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -597,28 +597,28 @@ class Observation(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.Observation """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field( + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceEndDate", "occurrenceEndDate"), serialization_alias="OccurrenceEndDate", ) - observation_source_concept: Optional[int] = Field( + observation_source_concept: int | None = Field( default=None, validation_alias=AliasChoices("ObservationSourceConcept", "observationSourceConcept"), serialization_alias="ObservationSourceConcept", ) - gender_cs: Optional[ConceptSetSelection] = Field( + gender_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("GenderCS", "genderCS"), serialization_alias="GenderCS", ) - observation_type: Optional[list[Concept]] = Field( + observation_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ObservationType", "observationType"), serialization_alias="ObservationType", ) - observation_type_cs: Optional[ConceptSetSelection] = Field( + observation_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ObservationTypeCS", "observationTypeCS"), serialization_alias="ObservationTypeCS", @@ -628,78 +628,78 @@ class Observation(Criteria): validation_alias=AliasChoices("ObservationTypeExclude", "observationTypeExclude"), serialization_alias="ObservationTypeExclude", ) - provider_specialty_cs: Optional[ConceptSetSelection] = Field( + provider_specialty_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialtyCS", "providerSpecialtyCS"), serialization_alias="ProviderSpecialtyCS", ) - visit_type_cs: Optional[ConceptSetSelection] = Field( + visit_type_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("VisitTypeCS", "visitTypeCS"), serialization_alias="VisitTypeCS", ) - visit_type: Optional[list[Concept]] = Field( + visit_type: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("VisitType", "visitType"), serialization_alias="VisitType", ) - value_as_number: Optional[NumericRange] = Field( + value_as_number: NumericRange | None = Field( default=None, validation_alias=AliasChoices("ValueAsNumber", "valueAsNumber"), serialization_alias="ValueAsNumber", ) - unit: Optional[list[Concept]] = Field( + unit: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("Unit", "unit"), serialization_alias="Unit", ) - unit_cs: Optional[ConceptSetSelection] = Field( + unit_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("UnitCS", "unitCS"), serialization_alias="UnitCS", ) - value_as_concept: Optional[list[Concept]] = Field( + value_as_concept: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ValueAsConcept", "valueAsConcept"), serialization_alias="ValueAsConcept", ) - value_as_concept_cs: Optional[ConceptSetSelection] = Field( + value_as_concept_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ValueAsConceptCS", "valueAsConceptCS"), serialization_alias="ValueAsConceptCS", ) - qualifier: Optional[list[Concept]] = Field( + qualifier: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("Qualifier", "qualifier"), serialization_alias="Qualifier", ) - qualifier_cs: Optional[ConceptSetSelection] = Field( + qualifier_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("QualifierCS", "qualifierCS"), serialization_alias="QualifierCS", ) - value_as_string: Optional[TextFilter] = Field( + value_as_string: TextFilter | None = Field( default=None, validation_alias=AliasChoices("ValueAsString", "valueAsString"), serialization_alias="ValueAsString", ) - codeset_id: Optional[int] = Field( + codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("CodesetId", "codesetId"), serialization_alias="CodesetId", ) - first: Optional[bool] = Field( + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - provider_specialty: Optional[list[Concept]] = Field( + provider_specialty: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ProviderSpecialty", "providerSpecialty"), serialization_alias="ProviderSpecialty", ) - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field( + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceStartDate", "occurrenceStartDate"), serialization_alias="OccurrenceStartDate", @@ -714,72 +714,72 @@ class Measurement(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.Measurement """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - measurement_source_concept: Optional[int] = Field(default=None, alias="MeasurementSourceConcept") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - measurement_type: Optional[list[Concept]] = Field(default=None, alias="MeasurementType") - measurement_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="MeasurementTypeCS") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + measurement_source_concept: int | None = Field(default=None, alias="MeasurementSourceConcept") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + measurement_type: list[Concept] | None = Field(default=None, alias="MeasurementType") + measurement_type_cs: ConceptSetSelection | None = Field(default=None, alias="MeasurementTypeCS") measurement_type_exclude: bool = Field( default=False, validation_alias=AliasChoices("MeasurementTypeExclude", "measurementTypeExclude"), serialization_alias="MeasurementTypeExclude", ) - operator: Optional[list[Concept]] = None - operator_cs: Optional[ConceptSetSelection] = Field(default=None, alias="OperatorCS") - value_as_number: Optional[NumericRange] = Field(default=None, alias="ValueAsNumber") - value_as_string: Optional[TextFilter] = Field(default=None, alias="ValueAsString") - unit: Optional[list[Concept]] = Field(default=None, alias="Unit") - unit_cs: Optional[ConceptSetSelection] = Field(default=None, alias="UnitCS") - range_low: Optional[NumericRange] = Field(default=None, alias="RangeLow") - range_high: Optional[NumericRange] = Field(default=None, alias="RangeHigh") - provider_specialty_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProviderSpecialtyCS") - visit_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="VisitTypeCS") - visit_type: Optional[list[Concept]] = Field(default=None, alias="VisitType") - codeset_id: Optional[int] = Field( + operator: list[Concept] | None = None + operator_cs: ConceptSetSelection | None = Field(default=None, alias="OperatorCS") + value_as_number: NumericRange | None = Field(default=None, alias="ValueAsNumber") + value_as_string: TextFilter | None = Field(default=None, alias="ValueAsString") + unit: list[Concept] | None = Field(default=None, alias="Unit") + unit_cs: ConceptSetSelection | None = Field(default=None, alias="UnitCS") + range_low: NumericRange | None = Field(default=None, alias="RangeLow") + range_high: NumericRange | None = Field(default=None, alias="RangeHigh") + provider_specialty_cs: ConceptSetSelection | None = Field(default=None, alias="ProviderSpecialtyCS") + visit_type_cs: ConceptSetSelection | None = Field(default=None, alias="VisitTypeCS") + visit_type: list[Concept] | None = Field(default=None, alias="VisitType") + codeset_id: int | None = Field( default=None, validation_alias=AliasChoices("CodesetId", "codesetId"), serialization_alias="CodesetId", ) - value_as_concept: Optional[list[Concept]] = Field( + value_as_concept: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ValueAsConcept", "valueAsConcept"), serialization_alias="ValueAsConcept", ) - value_as_concept_cs: Optional[ConceptSetSelection] = Field( + value_as_concept_cs: ConceptSetSelection | None = Field( default=None, validation_alias=AliasChoices("ValueAsConceptCS", "valueAsConceptCS"), serialization_alias="ValueAsConceptCS", ) - abnormal: Optional[bool] = Field( + abnormal: bool | None = Field( default=None, validation_alias=AliasChoices("Abnormal", "abnormal"), serialization_alias="Abnormal", ) - range_low_ratio: Optional[NumericRange] = Field( + range_low_ratio: NumericRange | None = Field( default=None, validation_alias=AliasChoices("RangeLowRatio", "rangeLowRatio"), serialization_alias="RangeLowRatio", ) - range_high_ratio: Optional[NumericRange] = Field( + range_high_ratio: NumericRange | None = Field( default=None, validation_alias=AliasChoices("RangeHighRatio", "rangeHighRatio"), serialization_alias="RangeHighRatio", ) - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") - visits: Optional[list[Concept]] = None # Placeholder if needed, but not in list - visit_type: Optional[list[Concept]] = Field(default=None, alias="VisitType") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") + visits: list[Concept] | None = None # Placeholder if needed, but not in list + visit_type: list[Concept] | None = Field(default=None, alias="VisitType") - first: Optional[bool] = Field( + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -790,27 +790,27 @@ class DeviceExposure(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.DeviceExposure """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - device_source_concept: Optional[int] = Field(default=None, alias="DeviceSourceConcept") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - device_type: Optional[list[Concept]] = Field(default=None, alias="DeviceType") - device_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="DeviceTypeCS") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + device_source_concept: int | None = Field(default=None, alias="DeviceSourceConcept") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + device_type: list[Concept] | None = Field(default=None, alias="DeviceType") + device_type_cs: ConceptSetSelection | None = Field(default=None, alias="DeviceTypeCS") device_type_exclude: bool = Field(default=False, alias="DeviceTypeExclude") - unique_device_id: Optional[TextFilter] = Field(default=None, alias="UniqueDeviceId") - quantity: Optional[NumericRange] = None - provider_specialty_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProviderSpecialtyCS") - visit_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="VisitTypeCS") - visit_type: Optional[list[Concept]] = Field(default=None, alias="VisitType") - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field( + unique_device_id: TextFilter | None = Field(default=None, alias="UniqueDeviceId") + quantity: NumericRange | None = None + provider_specialty_cs: ConceptSetSelection | None = Field(default=None, alias="ProviderSpecialtyCS") + visit_type_cs: ConceptSetSelection | None = Field(default=None, alias="VisitTypeCS") + visit_type: list[Concept] | None = Field(default=None, alias="VisitType") + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - age: Optional[NumericRange] = Field(default=None, alias="Age") - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + age: NumericRange | None = Field(default=None, alias="Age") + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -821,29 +821,29 @@ class Specimen(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.Specimen """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - specimen_source_concept: Optional[int] = Field(default=None, alias="SpecimenSourceConcept") - source_id: Optional[TextFilter] = Field(default=None, alias="SourceId") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - specimen_type: Optional[list[Concept]] = Field(default=None, alias="SpecimenType") - specimen_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="SpecimenTypeCS") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + specimen_source_concept: int | None = Field(default=None, alias="SpecimenSourceConcept") + source_id: TextFilter | None = Field(default=None, alias="SourceId") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + specimen_type: list[Concept] | None = Field(default=None, alias="SpecimenType") + specimen_type_cs: ConceptSetSelection | None = Field(default=None, alias="SpecimenTypeCS") specimen_type_exclude: bool = Field(default=False, alias="SpecimenTypeExclude") - unit: Optional[list[Concept]] = None - unit_cs: Optional[ConceptSetSelection] = Field(default=None, alias="UnitCS") - anatomic_site: Optional[list[Concept]] = Field(default=None, alias="AnatomicSite") - anatomic_site_cs: Optional[ConceptSetSelection] = Field(default=None, alias="AnatomicSiteCS") - disease_status: Optional[list[Concept]] = Field(default=None, alias="DiseaseStatus") - disease_status_cs: Optional[ConceptSetSelection] = Field(default=None, alias="DiseaseStatusCS") - quantity: Optional[NumericRange] = None - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field( + unit: list[Concept] | None = None + unit_cs: ConceptSetSelection | None = Field(default=None, alias="UnitCS") + anatomic_site: list[Concept] | None = Field(default=None, alias="AnatomicSite") + anatomic_site_cs: ConceptSetSelection | None = Field(default=None, alias="AnatomicSiteCS") + disease_status: list[Concept] | None = Field(default=None, alias="DiseaseStatus") + disease_status_cs: ConceptSetSelection | None = Field(default=None, alias="DiseaseStatusCS") + quantity: NumericRange | None = None + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -854,23 +854,23 @@ class Death(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.Death """ - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - occurrence_end_date: Optional[DateRange] = Field(default=None, alias="OccurrenceEndDate") - death_source_concept: Optional[int] = Field(default=None, alias="DeathSourceConcept") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - death_type: Optional[list[Concept]] = Field(default=None, alias="DeathType") - death_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="DeathTypeCS") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + occurrence_end_date: DateRange | None = Field(default=None, alias="OccurrenceEndDate") + death_source_concept: int | None = Field(default=None, alias="DeathSourceConcept") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + death_type: list[Concept] | None = Field(default=None, alias="DeathType") + death_type_cs: ConceptSetSelection | None = Field(default=None, alias="DeathTypeCS") death_type_exclude: bool = Field( default=False, validation_alias=AliasChoices("DeathTypeExclude", "deathTypeExclude"), serialization_alias="DeathTypeExclude", ) - cause_source_concept: Optional[int] = Field(default=None, alias="CauseSourceConcept") - cause_source_concept_cs: Optional[ConceptSetSelection] = Field(default=None, alias="CauseSourceConceptCS") - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") + cause_source_concept: int | None = Field(default=None, alias="CauseSourceConcept") + cause_source_concept_cs: ConceptSetSelection | None = Field(default=None, alias="CauseSourceConceptCS") + codeset_id: int | None = Field(default=None, alias="CodesetId") - age: Optional[NumericRange] = None - occurrence_start_date: Optional[DateRange] = Field(default=None, alias="OccurrenceStartDate") + age: NumericRange | None = None + occurrence_start_date: DateRange | None = Field(default=None, alias="OccurrenceStartDate") model_config = ConfigDict(populate_by_name=True) @@ -881,25 +881,25 @@ class VisitDetail(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.VisitDetail """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field(default=None, alias="First") - visit_detail_start_date: Optional[DateRange] = Field(default=None, alias="VisitDetailStartDate") - visit_detail_end_date: Optional[DateRange] = Field(default=None, alias="VisitDetailEndDate") - visit_detail_type: Optional[list[Concept]] = Field(default=None, alias="VisitDetailType") - visit_detail_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="VisitDetailTypeCS") + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field(default=None, alias="First") + visit_detail_start_date: DateRange | None = Field(default=None, alias="VisitDetailStartDate") + visit_detail_end_date: DateRange | None = Field(default=None, alias="VisitDetailEndDate") + visit_detail_type: list[Concept] | None = Field(default=None, alias="VisitDetailType") + visit_detail_type_cs: ConceptSetSelection | None = Field(default=None, alias="VisitDetailTypeCS") visit_detail_type_exclude: bool = Field(default=False, alias="VisitDetailTypeExclude") - visit_detail_source_concept: Optional[int] = Field(default=None, alias="VisitDetailSourceConcept") - visit_detail_length: Optional[NumericRange] = Field(default=None, alias="VisitDetailLength") - age: Optional[NumericRange] = Field(default=None, alias="Age") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - provider_specialty: Optional[list[Concept]] = Field(default=None, alias="ProviderSpecialty") - provider_specialty_cs: Optional[ConceptSetSelection] = Field(default=None, alias="ProviderSpecialtyCS") - place_of_service: Optional[list[Concept]] = Field(default=None, alias="PlaceOfService") - place_of_service_cs: Optional[ConceptSetSelection] = Field(default=None, alias="PlaceOfServiceCS") - place_of_service_location: Optional[int] = Field(default=None, alias="PlaceOfServiceLocation") - discharge_to: Optional[list[Concept]] = Field(default=None, alias="DischargeTo") - discharge_to_cs: Optional[ConceptSetSelection] = Field(default=None, alias="DischargeToCS") + visit_detail_source_concept: int | None = Field(default=None, alias="VisitDetailSourceConcept") + visit_detail_length: NumericRange | None = Field(default=None, alias="VisitDetailLength") + age: NumericRange | None = Field(default=None, alias="Age") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + provider_specialty: list[Concept] | None = Field(default=None, alias="ProviderSpecialty") + provider_specialty_cs: ConceptSetSelection | None = Field(default=None, alias="ProviderSpecialtyCS") + place_of_service: list[Concept] | None = Field(default=None, alias="PlaceOfService") + place_of_service_cs: ConceptSetSelection | None = Field(default=None, alias="PlaceOfServiceCS") + place_of_service_location: int | None = Field(default=None, alias="PlaceOfServiceLocation") + discharge_to: list[Concept] | None = Field(default=None, alias="DischargeTo") + discharge_to_cs: ConceptSetSelection | None = Field(default=None, alias="DischargeToCS") model_config = ConfigDict(populate_by_name=True) @@ -910,15 +910,15 @@ class ObservationPeriod(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.ObservationPeriod """ - first: Optional[bool] = Field(default=None, alias="First") - period_start_date: Optional[DateRange] = Field(default=None, alias="PeriodStartDate") - period_end_date: Optional[DateRange] = Field(default=None, alias="PeriodEndDate") - user_defined_period: Optional[Period] = Field(default=None, alias="UserDefinedPeriod") - period_type: Optional[list[Concept]] = Field(default=None, alias="PeriodType") - period_type_cs: Optional[ConceptSetSelection] = Field(default=None, alias="PeriodTypeCS") - period_length: Optional[NumericRange] = Field(default=None, alias="PeriodLength") - age_at_start: Optional[NumericRange] = Field(default=None, alias="AgeAtStart") - age_at_end: Optional[NumericRange] = Field(default=None, alias="AgeAtEnd") + first: bool | None = Field(default=None, alias="First") + period_start_date: DateRange | None = Field(default=None, alias="PeriodStartDate") + period_end_date: DateRange | None = Field(default=None, alias="PeriodEndDate") + user_defined_period: Period | None = Field(default=None, alias="UserDefinedPeriod") + period_type: list[Concept] | None = Field(default=None, alias="PeriodType") + period_type_cs: ConceptSetSelection | None = Field(default=None, alias="PeriodTypeCS") + period_length: NumericRange | None = Field(default=None, alias="PeriodLength") + age_at_start: NumericRange | None = Field(default=None, alias="AgeAtStart") + age_at_end: NumericRange | None = Field(default=None, alias="AgeAtEnd") model_config = ConfigDict(populate_by_name=True) @@ -929,23 +929,23 @@ class PayerPlanPeriod(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.PayerPlanPeriod """ - first: Optional[bool] = Field(default=None, alias="First") - period_start_date: Optional[DateRange] = Field(default=None, alias="PeriodStartDate") - period_end_date: Optional[DateRange] = Field(default=None, alias="PeriodEndDate") - user_defined_period: Optional[Period] = Field(default=None, alias="UserDefinedPeriod") - period_length: Optional[NumericRange] = Field(default=None, alias="PeriodLength") - age_at_start: Optional[NumericRange] = Field(default=None, alias="AgeAtStart") - age_at_end: Optional[NumericRange] = Field(default=None, alias="AgeAtEnd") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - payer_concept: Optional[int] = Field(default=None, alias="PayerConcept") - plan_concept: Optional[int] = Field(default=None, alias="PlanConcept") - sponsor_concept: Optional[int] = Field(default=None, alias="SponsorConcept") - stop_reason_concept: Optional[int] = Field(default=None, alias="StopReasonConcept") - payer_source_concept: Optional[int] = Field(default=None, alias="PayerSourceConcept") - plan_source_concept: Optional[int] = Field(default=None, alias="PlanSourceConcept") - sponsor_source_concept: Optional[int] = Field(default=None, alias="SponsorSourceConcept") - stop_reason_source_concept: Optional[int] = Field(default=None, alias="StopReasonSourceConcept") + first: bool | None = Field(default=None, alias="First") + period_start_date: DateRange | None = Field(default=None, alias="PeriodStartDate") + period_end_date: DateRange | None = Field(default=None, alias="PeriodEndDate") + user_defined_period: Period | None = Field(default=None, alias="UserDefinedPeriod") + period_length: NumericRange | None = Field(default=None, alias="PeriodLength") + age_at_start: NumericRange | None = Field(default=None, alias="AgeAtStart") + age_at_end: NumericRange | None = Field(default=None, alias="AgeAtEnd") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + payer_concept: int | None = Field(default=None, alias="PayerConcept") + plan_concept: int | None = Field(default=None, alias="PlanConcept") + sponsor_concept: int | None = Field(default=None, alias="SponsorConcept") + stop_reason_concept: int | None = Field(default=None, alias="StopReasonConcept") + payer_source_concept: int | None = Field(default=None, alias="PayerSourceConcept") + plan_source_concept: int | None = Field(default=None, alias="PlanSourceConcept") + sponsor_source_concept: int | None = Field(default=None, alias="SponsorSourceConcept") + stop_reason_source_concept: int | None = Field(default=None, alias="StopReasonSourceConcept") model_config = ConfigDict(populate_by_name=True) @@ -956,7 +956,7 @@ class LocationRegion(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.LocationRegion """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") + codeset_id: int | None = Field(default=None, alias="CodesetId") model_config = ConfigDict(populate_by_name=True) @@ -972,21 +972,21 @@ class ConditionEra(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.ConditionEra """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field( + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - era_start_date: Optional[DateRange] = Field(default=None, alias="EraStartDate") - era_end_date: Optional[DateRange] = Field(default=None, alias="EraEndDate") - occurrence_count: Optional[NumericRange] = Field(default=None, alias="OccurrenceCount") - era_length: Optional[NumericRange] = Field(default=None, alias="EraLength") - age_at_start: Optional[NumericRange] = Field(default=None, alias="AgeAtStart") - age_at_end: Optional[NumericRange] = Field(default=None, alias="AgeAtEnd") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - date_adjustment: Optional[DateAdjustment] = Field(default=None, alias="DateAdjustment") + era_start_date: DateRange | None = Field(default=None, alias="EraStartDate") + era_end_date: DateRange | None = Field(default=None, alias="EraEndDate") + occurrence_count: NumericRange | None = Field(default=None, alias="OccurrenceCount") + era_length: NumericRange | None = Field(default=None, alias="EraLength") + age_at_start: NumericRange | None = Field(default=None, alias="AgeAtStart") + age_at_end: NumericRange | None = Field(default=None, alias="AgeAtEnd") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + date_adjustment: DateAdjustment | None = Field(default=None, alias="DateAdjustment") model_config = ConfigDict(populate_by_name=True) @@ -997,22 +997,22 @@ class DrugEra(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.DrugEra """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field( + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field( default=None, validation_alias=AliasChoices("First", "first"), serialization_alias="First", ) - era_start_date: Optional[DateRange] = Field(default=None, alias="EraStartDate") - era_end_date: Optional[DateRange] = Field(default=None, alias="EraEndDate") - occurrence_count: Optional[NumericRange] = Field(default=None, alias="OccurrenceCount") - gap_days: Optional[NumericRange] = Field(default=None, alias="GapDays") - era_length: Optional[NumericRange] = Field(default=None, alias="EraLength") - age_at_start: Optional[NumericRange] = Field(default=None, alias="AgeAtStart") - age_at_end: Optional[NumericRange] = Field(default=None, alias="AgeAtEnd") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") - date_adjustment: Optional[DateAdjustment] = Field(default=None, alias="DateAdjustment") + era_start_date: DateRange | None = Field(default=None, alias="EraStartDate") + era_end_date: DateRange | None = Field(default=None, alias="EraEndDate") + occurrence_count: NumericRange | None = Field(default=None, alias="OccurrenceCount") + gap_days: NumericRange | None = Field(default=None, alias="GapDays") + era_length: NumericRange | None = Field(default=None, alias="EraLength") + age_at_start: NumericRange | None = Field(default=None, alias="AgeAtStart") + age_at_end: NumericRange | None = Field(default=None, alias="AgeAtEnd") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") + date_adjustment: DateAdjustment | None = Field(default=None, alias="DateAdjustment") model_config = ConfigDict(populate_by_name=True) @@ -1023,18 +1023,18 @@ class DoseEra(Criteria): Java equivalent: org.ohdsi.circe.cohortdefinition.DoseEra """ - codeset_id: Optional[int] = Field(default=None, alias="CodesetId") - first: Optional[bool] = Field(default=None, alias="First") - era_start_date: Optional[DateRange] = Field(default=None, alias="EraStartDate") - era_end_date: Optional[DateRange] = Field(default=None, alias="EraEndDate") - unit: Optional[list[Concept]] = Field(default=None, alias="Unit") - unit_cs: Optional[ConceptSetSelection] = Field(default=None, alias="UnitCS") - dose_value: Optional[NumericRange] = Field(default=None, alias="DoseValue") - era_length: Optional[NumericRange] = Field(default=None, alias="EraLength") - age_at_start: Optional[NumericRange] = Field(default=None, alias="AgeAtStart") - age_at_end: Optional[NumericRange] = Field(default=None, alias="AgeAtEnd") - gender: Optional[list[Concept]] = Field(default=None, serialization_alias="gender") - gender_cs: Optional[ConceptSetSelection] = Field(default=None, alias="GenderCS") + codeset_id: int | None = Field(default=None, alias="CodesetId") + first: bool | None = Field(default=None, alias="First") + era_start_date: DateRange | None = Field(default=None, alias="EraStartDate") + era_end_date: DateRange | None = Field(default=None, alias="EraEndDate") + unit: list[Concept] | None = Field(default=None, alias="Unit") + unit_cs: ConceptSetSelection | None = Field(default=None, alias="UnitCS") + dose_value: NumericRange | None = Field(default=None, alias="DoseValue") + era_length: NumericRange | None = Field(default=None, alias="EraLength") + age_at_start: NumericRange | None = Field(default=None, alias="AgeAtStart") + age_at_end: NumericRange | None = Field(default=None, alias="AgeAtEnd") + gender: list[Concept] | None = Field(default=None, serialization_alias="gender") + gender_cs: ConceptSetSelection | None = Field(default=None, alias="GenderCS") model_config = ConfigDict(populate_by_name=True) @@ -1069,7 +1069,7 @@ class CriteriaGroup(BaseModel): validation_alias=AliasChoices("CriteriaList", "criteriaList"), serialization_alias="CriteriaList", ) - count: Optional[int] = Field( + count: int | None = Field( default=None, validation_alias=AliasChoices("Count", "count"), serialization_alias="Count", @@ -1084,7 +1084,7 @@ class CriteriaGroup(BaseModel): validation_alias=AliasChoices("DemographicCriteriaList", "demographicCriteriaList"), serialization_alias="DemographicCriteriaList", ) - type: Optional[str] = Field( + type: str | None = Field( default=None, validation_alias=AliasChoices("Type", "type"), serialization_alias="Type", @@ -1372,25 +1372,25 @@ def normalize_window(window_dict: dict) -> dict: # Define CriteriaType Union for strict typing. # Criteria is last so known subtypes are tried first; it also acts as # a catch-all that accepts any registered extension subclass. -_CriteriaTypeUnion = Union[ - ConditionOccurrence, - DrugExposure, - ProcedureOccurrence, - VisitOccurrence, - Observation, - Measurement, - DeviceExposure, - Specimen, - Death, - VisitDetail, - ObservationPeriod, - PayerPlanPeriod, - LocationRegion, - ConditionEra, - DrugEra, - DoseEra, - Criteria, # catch-all for extension subclasses -] +_CriteriaTypeUnion = ( + ConditionOccurrence + | DrugExposure + | ProcedureOccurrence + | VisitOccurrence + | Observation + | Measurement + | DeviceExposure + | Specimen + | Death + | VisitDetail + | ObservationPeriod + | PayerPlanPeriod + | LocationRegion + | ConditionEra + | DrugEra + | DoseEra + | Criteria # catch-all for extension subclasses +) def _validate_criteria_extension(v: Any) -> Any: @@ -1443,12 +1443,12 @@ class PrimaryCriteria(BaseModel): validation_alias=AliasChoices("CriteriaList", "criteriaList"), serialization_alias="CriteriaList", ) - observation_window: Optional[ObservationFilter] = Field( + observation_window: ObservationFilter | None = Field( default=None, validation_alias=AliasChoices("ObservationWindow", "observationWindow"), serialization_alias="ObservationWindow", ) - primary_limit: Optional[ResultLimit] = Field( + primary_limit: ResultLimit | None = Field( default=None, validation_alias=AliasChoices( "PrimaryLimit", diff --git a/circe/cohortdefinition/interfaces.py b/circe/cohortdefinition/interfaces.py index 0492ab66..65d7ffa8 100644 --- a/circe/cohortdefinition/interfaces.py +++ b/circe/cohortdefinition/interfaces.py @@ -10,7 +10,6 @@ """ from abc import ABC, abstractmethod -from typing import Optional, Union from .builders.utils import BuilderOptions from .core import CustomEraStrategy, DateOffsetStrategy @@ -34,24 +33,24 @@ ) # Type alias for all criteria types -Criteria = Union[ - LocationRegion, - ConditionEra, - ConditionOccurrence, - Death, - DeviceExposure, - DoseEra, - DrugEra, - DrugExposure, - Measurement, - Observation, - ObservationPeriod, - PayerPlanPeriod, - ProcedureOccurrence, - Specimen, - VisitOccurrence, - VisitDetail, -] +Criteria = ( + LocationRegion + | ConditionEra + | ConditionOccurrence + | Death + | DeviceExposure + | DoseEra + | DrugEra + | DrugExposure + | Measurement + | Observation + | ObservationPeriod + | PayerPlanPeriod + | ProcedureOccurrence + | Specimen + | VisitOccurrence + | VisitDetail +) class IGetCriteriaSqlDispatcher(ABC): @@ -61,7 +60,7 @@ class IGetCriteriaSqlDispatcher(ABC): """ @abstractmethod - def get_criteria_sql(self, criteria: Criteria, options: Optional[BuilderOptions] = None) -> str: + def get_criteria_sql(self, criteria: Criteria, options: BuilderOptions | None = None) -> str: """Generate SQL for various criteria types. Args: @@ -75,7 +74,7 @@ def get_criteria_sql(self, criteria: Criteria, options: Optional[BuilderOptions] # Type alias for end strategies -EndStrategy = Union[DateOffsetStrategy, CustomEraStrategy] +EndStrategy = DateOffsetStrategy | CustomEraStrategy class IGetEndStrategySqlDispatcher(ABC): diff --git a/circe/cohortdefinition/printfriendly/markdown_render.py b/circe/cohortdefinition/printfriendly/markdown_render.py index e2a1c761..48423ede 100644 --- a/circe/cohortdefinition/printfriendly/markdown_render.py +++ b/circe/cohortdefinition/printfriendly/markdown_render.py @@ -14,7 +14,6 @@ import json from datetime import datetime from pathlib import Path -from typing import Optional, Union import jinja2 @@ -34,9 +33,9 @@ class MarkdownRender: def __init__( self, - concept_sets: Optional[list[ConceptSet]] = None, + concept_sets: list[ConceptSet] | None = None, include_concept_sets: bool = False, - template_paths: Optional[list[Path]] = None, + template_paths: list[Path] | None = None, ): """Initialize the markdown renderer. @@ -90,9 +89,9 @@ def get_template_for_criteria(criteria): def render_cohort_expression( self, - cohort_expression: Union[CohortExpression, str], - include_concept_sets: Optional[bool] = None, - title: Optional[str] = None, + cohort_expression: CohortExpression | str, + include_concept_sets: bool | None = None, + title: str | None = None, ) -> str: """Render a cohort expression to markdown format. @@ -133,7 +132,7 @@ def render_cohort_expression( include_concept_sets=should_include, ) - def render_concept_set_list(self, concept_sets: Union[list[ConceptSet], str]) -> str: + def render_concept_set_list(self, concept_sets: list[ConceptSet] | str) -> str: """Render a list of concept sets to markdown format. Java equivalent: renderConceptSetList(ConceptSet[]) @@ -164,7 +163,7 @@ def render_concept_set_list(self, concept_sets: Union[list[ConceptSet], str]) -> return template.render(conceptSets=concept_sets) - def render_concept_set(self, concept_set: Union[ConceptSet, str]) -> str: + def render_concept_set(self, concept_set: ConceptSet | str) -> str: """Render a single concept set to markdown format. Java equivalent: renderConceptSet(ConceptSet) @@ -186,7 +185,7 @@ def render_concept_set(self, concept_set: Union[ConceptSet, str]) -> str: # Custom Filters and Functions (matching Java utils.ftl) # ========================================================================= - def _codeset_name(self, codeset_id: Optional[int], default_name: str = "any") -> str: + def _codeset_name(self, codeset_id: int | None, default_name: str = "any") -> str: """Get concept set name from codeset ID, or return default. Java equivalent: utils.codesetName() @@ -228,7 +227,7 @@ def _format_date(self, date_string: str) -> str: except (ValueError, AttributeError): return "_invalid date_" - def _format_number(self, value: Union[int, float]) -> str: + def _format_number(self, value: int | float) -> str: """Format number with thousands separators and handle integer/float logic. Args: diff --git a/circe/execution/__init__.py b/circe/execution/__init__.py index 180d7927..ba27df0a 100644 --- a/circe/execution/__init__.py +++ b/circe/execution/__init__.py @@ -13,12 +13,10 @@ UnsupportedCriterionError, UnsupportedFeatureError, ) -from .ibis.codesets import clear_codeset_cache __all__ = [ "build_cohort", "write_cohort", - "clear_codeset_cache", "apply_databricks_post_connect_workaround", "ExecutionError", "ExecutionNormalizationError", diff --git a/circe/execution/_dataclass.py b/circe/execution/_dataclass.py index f7129f39..f2b4f35c 100644 --- a/circe/execution/_dataclass.py +++ b/circe/execution/_dataclass.py @@ -1,8 +1,8 @@ from __future__ import annotations -import sys +from collections.abc import Callable from dataclasses import dataclass -from typing import Any, Callable, TypeVar, cast, overload +from typing import Any, TypeVar, cast, overload from typing_extensions import dataclass_transform @@ -29,10 +29,7 @@ def frozen_slots_dataclass( """ def wrap(cls: type[T]) -> type[T]: - dataclass_factory = cast(Any, dataclass) - if sys.version_info >= (3, 10): - return cast(type[T], dataclass_factory(frozen=True, slots=True, **kwargs)(cls)) - return cast(type[T], dataclass_factory(frozen=True, **kwargs)(cls)) + return cast(type[T], dataclass(frozen=True, slots=True, **kwargs)(cls)) if _cls is None: return wrap diff --git a/circe/execution/api.py b/circe/execution/api.py index a9574b87..5401c1d9 100644 --- a/circe/execution/api.py +++ b/circe/execution/api.py @@ -6,6 +6,7 @@ from .databricks_compat import maybe_apply_databricks_post_connect_workaround from .engine.cohort import build_cohort_table from .errors import ExecutionError +from .ibis.codesets import build_single_codeset_table from .ibis.context import make_execution_context from .ibis.materialize import project_to_ohdsi_cohort_table from .ibis.operations import ( @@ -29,23 +30,51 @@ def build_cohort( cdm_schema: str, results_schema: str | None = None, vocabulary_schema: str | None = None, - use_persistent_cache: bool = False, + cohort_id: int = 0, + materialize: bool = True, + codeset_table: Table | None = None, + cohort_table: str = "cohort", + session_prefix: str = "", ) -> Table: - """Normalize, compile, and assemble a cohort relation.""" + """Normalize, compile, and assemble a cohort relation. + + Paths through stage-by-stage temp tables when *cohort_id* is provided + and *materialize* is True, so that the ibis expression tree never grows + too large to compile. Set *materialize=False* for compile-only use. + + When *codeset_table* is provided it is used directly. Otherwise a + per-cohort codeset table is auto-created. + """ maybe_apply_databricks_post_connect_workaround(backend) normalized = normalize_cohort(expression) + if codeset_table is None: + codeset_table = build_single_codeset_table( + backend=backend, + concept_sets=normalized.concept_sets, + batch_table_name="__codesets", + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + session_prefix=session_prefix, + ) + ctx = make_execution_context( backend=backend, cdm_schema=cdm_schema, results_schema=results_schema, vocabulary_schema=vocabulary_schema, - concept_sets=normalized.concept_sets, - use_persistent_cache=use_persistent_cache, + codeset_table=codeset_table, ) - return build_cohort_table(normalized, ctx) + return build_cohort_table( + normalized, + ctx, + cohort_id=cohort_id, + materialize=materialize, + cohort_table=cohort_table, + session_prefix=session_prefix, + ) def write_relation( @@ -87,8 +116,9 @@ def write_relation( def write_cohort( - expression: CohortExpression, + expression: CohortExpression | None = None, *, + compiled_relation: Table | None = None, backend: IbisBackendLike, cdm_schema: str, cohort_table: str, @@ -96,21 +126,50 @@ def write_cohort( results_schema: str | None = None, vocabulary_schema: str | None = None, if_exists: Literal["fail", "replace"] = "fail", - use_persistent_cache: bool = False, ) -> None: - """Build cohort rows and materialize them with cohort-scoped semantics.""" + """Build cohort rows and materialize them with cohort-scoped semantics. + + Args: + expression: Cohort expression to compile and execute. Provide one of + ``expression`` or ``compiled_relation`` (not both). + compiled_relation: A pre-compiled ibis relation (output of + ``build_cohort()`` projected with ``project_to_ohdsi_cohort_table()``). + When provided, the compilation step is skipped and this relation is + materialized directly. Use this to isolate database-execution time + from query-compilation time in benchmarks. + backend: Ibis backend connection. + cdm_schema: Schema containing the OMOP CDM source tables. + cohort_table: Name of the OHDSI cohort table to write results into. + cohort_id: The cohort_definition_id value to stamp on written rows. + results_schema: Schema for the cohort table. + vocabulary_schema: Schema for vocabulary tables (defaults to cdm_schema). + if_exists: Behaviour when cohort rows already exist. One of + ``"fail"`` (raise) or ``"replace"`` (remove existing rows for + this cohort_id before writing). + use_persistent_cache: Whether to cache concept set lookups persistently. + + Raises: + ValueError: If both or neither of ``expression`` / ``compiled_relation`` + are provided, or ``if_exists`` is invalid. + ExecutionError: If the write fails. + """ + if (expression is None) == (compiled_relation is None): + raise ValueError("Exactly one of expression or compiled_relation must be provided.") if if_exists not in {"fail", "replace"}: raise ValueError("if_exists must be one of {'fail', 'replace'} for write_cohort.") - new_rows = build_cohort( - expression, - backend=backend, - cdm_schema=cdm_schema, - results_schema=results_schema, - vocabulary_schema=vocabulary_schema, - use_persistent_cache=use_persistent_cache, - ) - new_rows = project_to_ohdsi_cohort_table(new_rows, cohort_id=cohort_id) + if compiled_relation is not None: + new_rows = compiled_relation + else: + new_rows = build_cohort( + expression, # type: ignore[arg-type] + backend=backend, + cdm_schema=cdm_schema, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + cohort_id=cohort_id, + ) + new_rows = project_to_ohdsi_cohort_table(new_rows, cohort_id=cohort_id) if not table_exists(backend, table_name=cohort_table, schema=results_schema): write_relation( diff --git a/circe/execution/engine/censoring.py b/circe/execution/engine/censoring.py index ce21f626..264234dc 100644 --- a/circe/execution/engine/censoring.py +++ b/circe/execution/engine/censoring.py @@ -9,10 +9,12 @@ def _union_all(tables): - current = tables[0] - for table in tables[1:]: - current = current.union(table, distinct=False) - return current + if len(tables) == 1: + return tables[0] + mid = len(tables) // 2 + left = _union_all(tables[:mid]) + right = _union_all(tables[mid:]) + return left.union(right, distinct=False) def _compile_censor_events(criteria, ctx): diff --git a/circe/execution/engine/cohort.py b/circe/execution/engine/cohort.py index 7f34652a..4335bb37 100644 --- a/circe/execution/engine/cohort.py +++ b/circe/execution/engine/cohort.py @@ -1,6 +1,9 @@ from __future__ import annotations +import contextlib + from ..ibis.context import ExecutionContext +from ..ibis.operations import create_table, read_table from ..lower.criteria import lower_criterion from ..normalize.cohort import NormalizedCohort from ..plan.cohort import CohortPlan, PrimaryEventInput @@ -14,7 +17,51 @@ from .primary import build_primary_events -def build_cohort_table(normalized: NormalizedCohort, ctx: ExecutionContext) -> Table: +def _materialize( + table: Table, + *, + ctx: ExecutionContext, + cohort_id: int, + stage: str, + schema: str | None, + cohort_table: str = "cohort", + session_prefix: str = "", +) -> Table: + """Write *table* to a backend staging table and return a fresh reference. + + Uses a single session-scoped table per stage (not per cohort) since cohorts + are processed sequentially. The table is overwritten for each cohort. + """ + name = f"{session_prefix}__staging_{stage}" + create_table(ctx.backend, table_name=name, schema=schema, obj=table, overwrite=True) + return read_table(ctx.backend, table_name=name, schema=schema) + + +def _drop_staging_tables( + ctx: ExecutionContext, + schema: str | None, + session_prefix: str = "", +) -> None: + """Remove all session-scoped staging tables from the database.""" + for stage in ("primary", "qualified", "included", "ended"): + name = f"{session_prefix}__staging_{stage}" + with contextlib.suppress(Exception): + ctx.backend.drop_table(name, database=schema, force=True) + # Also drop the session-scoped codeset table + codeset_name = f"{session_prefix}__codesets" + with contextlib.suppress(Exception): + ctx.backend.drop_table(codeset_name, database=schema, force=True) + + +def build_cohort_table( + normalized: NormalizedCohort, + ctx: ExecutionContext, + *, + cohort_id: int = 0, + materialize: bool = True, + cohort_table: str = "cohort", + session_prefix: str = "", +) -> Table: primary_plans = tuple( PrimaryEventInput( event_plan=lower_criterion(criterion, criterion_index=index), @@ -29,27 +76,89 @@ def build_cohort_table(normalized: NormalizedCohort, ctx: ExecutionContext) -> T qualified_limit_type=normalized.result_limits.qualified_limit_type, expression_limit_type=normalized.result_limits.expression_limit_type, ) + + schema = ctx.results_schema or ctx.cdm_schema + + # ── Primary events ────────────────────────────────────────────────── primary_events = build_primary_events(cohort_plan, ctx) + has_additional_criteria = ( + normalized.additional_criteria is not None and not normalized.additional_criteria.is_empty() + ) + if materialize and has_additional_criteria: + primary_events = _materialize( + primary_events, + ctx=ctx, + cohort_id=cohort_id, + stage="primary", + schema=schema, + session_prefix=session_prefix, + cohort_table=cohort_table, + ) + + # ── Additional (correlated) criteria ──────────────────────────────── qualified_events = apply_additional_criteria(primary_events, normalized.additional_criteria, ctx) - if normalized.additional_criteria is not None and not normalized.additional_criteria.is_empty(): - qualified_events = apply_result_limit( + if has_additional_criteria: + qualified_events = apply_result_limit(qualified_events, cohort_plan.qualified_limit_type) + if materialize: + qualified_events = _materialize( qualified_events, - cohort_plan.qualified_limit_type, + ctx=ctx, + cohort_id=cohort_id, + stage="qualified", + schema=schema, + session_prefix=session_prefix, + cohort_table=cohort_table, ) - included_events = apply_inclusion_rules(qualified_events, normalized.inclusion_rules, ctx) - included_events = apply_result_limit( - included_events, - cohort_plan.expression_limit_type, - ) + + # ── Inclusion rules ───────────────────────────────────────────────── + # Materialise after every inclusion rule so that the ibis expression tree + # never grows deeper than one rule's worth of operations. Without this a + # cohort with N rules builds an N-level tree that, when compiled into a + # single SQL statement, produces query plans too large for some backends + # (e.g. Databricks Spark) to execute without resource exhaustion. + if materialize and normalized.inclusion_rules: + included_events = qualified_events + for rule in normalized.inclusion_rules: + included_events = apply_additional_criteria(included_events, rule.expression, ctx) + included_events = _materialize( + included_events, + ctx=ctx, + cohort_id=cohort_id, + stage="included", + schema=schema, + session_prefix=session_prefix, + cohort_table=cohort_table, + ) + included_events = apply_result_limit(included_events, cohort_plan.expression_limit_type) + else: + included_events = apply_inclusion_rules(qualified_events, normalized.inclusion_rules, ctx) + included_events = apply_result_limit(included_events, cohort_plan.expression_limit_type) + if materialize: + included_events = _materialize( + included_events, + ctx=ctx, + cohort_id=cohort_id, + stage="included", + schema=schema, + session_prefix=session_prefix, + cohort_table=cohort_table, + ) + + # ── End strategy ──────────────────────────────────────────────────── ended_events = apply_end_strategy(included_events, normalized.end_strategy, ctx) + if materialize: + ended_events = _materialize( + ended_events, + ctx=ctx, + cohort_id=cohort_id, + stage="ended", + schema=schema, + session_prefix=session_prefix, + cohort_table=cohort_table, + ) + + # ── Censoring + collapse (final stage — no materialize after) ────── censored_events = apply_censoring( - ended_events, - normalized.censoring_criteria, - normalized.censor_window, - ctx, - ) - return collapse_events( - censored_events, - normalized.collapse_settings, - normalized.censor_window, + ended_events, normalized.censoring_criteria, normalized.censor_window, ctx ) + return collapse_events(censored_events, normalized.collapse_settings, normalized.censor_window) diff --git a/circe/execution/engine/collapse.py b/circe/execution/engine/collapse.py index b6d0cc39..ed5114a2 100644 --- a/circe/execution/engine/collapse.py +++ b/circe/execution/engine/collapse.py @@ -2,7 +2,15 @@ import ibis -from ..plan.schema import END_DATE, PERSON_ID, START_DATE +from ..plan.schema import END_DATE, OP_END_DATE, OP_START_DATE, PERSON_ID, START_DATE + + +def _strip_op_columns(events): + """Remove internal observation period columns from final output.""" + cols_to_drop = [c for c in events.columns if c in (OP_START_DATE, OP_END_DATE)] + if cols_to_drop: + return events.drop(*cols_to_drop) + return events def _apply_censor_window(events, censor_window): @@ -71,11 +79,11 @@ def _collapse_era(intervals, era_pad: int): def collapse_events(events, collapse_settings, censor_window): if collapse_settings is None: - return _apply_censor_window(events, censor_window) + return _strip_op_columns(_apply_censor_window(events, censor_window)) collapse_type = (collapse_settings.collapse_type or "era").lower() if collapse_type == "no_collapse": - return _apply_censor_window(events, censor_window) + return _strip_op_columns(_apply_censor_window(events, censor_window)) intervals = events.select( events.person_id.cast("int64").name(PERSON_ID), diff --git a/circe/execution/engine/custom_era.py b/circe/execution/engine/custom_era.py new file mode 100644 index 00000000..23204802 --- /dev/null +++ b/circe/execution/engine/custom_era.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +import ibis + +from ..plan.schema import PERSON_ID, START_DATE +from .end_strategy import _replace_end_date, attach_observation_bounds + + +def _compute_exposure_end_date(table, *, days_supply_override: int | None): + start = table["drug_exposure_start_date"].cast("date") + + if days_supply_override is not None: + return start + ibis.interval(days=days_supply_override) + + raw_end = ( + table["drug_exposure_end_date"].cast("date") + if "drug_exposure_end_date" in table.columns + else ibis.null().cast("date") + ) + days_supply = ( + table["days_supply"].cast("int64") if "days_supply" in table.columns else ibis.null().cast("int64") + ) + supply_end = start + days_supply.as_interval("D") + + return ibis.coalesce(raw_end, supply_end, start + ibis.interval(days=1)) + + +def _compute_eras(exposures, *, gap_days: int, offset: int): + padded = exposures.mutate( + _padded_end=(exposures._exposure_end + ibis.interval(days=int(gap_days + offset))) + ) + + ordering = [ + padded.start_date, + padded._padded_end.desc(), + padded._exposure_end.desc(), + ] + + cumulative_window = ibis.cumulative_window(group_by=padded.person_id, order_by=ordering) + ordered_window = ibis.window(group_by=padded.person_id, order_by=ordering) + + with_cummax = padded.mutate(_cummax_padded_end=padded._padded_end.max().over(cumulative_window)) + + with_prev = with_cummax.mutate(_prev_max=with_cummax._cummax_padded_end.lag().over(ordered_window)) + + marked = with_prev.mutate( + _is_new=ibis.ifelse( + with_prev._prev_max.isnull() | (with_prev._prev_max < with_prev.start_date), + ibis.literal(1, type="int64"), + ibis.literal(0, type="int64"), + ) + ) + + group_window = ibis.cumulative_window( + group_by=marked.person_id, + order_by=[ + marked.start_date, + marked._padded_end.desc(), + marked._exposure_end.desc(), + marked._is_new.desc(), + ], + ) + era_indexed = marked.mutate(_era_id=marked._is_new.sum().over(group_window)) + + collapsed = era_indexed.group_by(era_indexed.person_id, era_indexed._era_id).aggregate( + era_start_date=era_indexed.start_date.min(), + _max_padded_end=era_indexed._padded_end.max(), + ) + + return collapsed.select( + collapsed.person_id.cast("int64").name(PERSON_ID), + collapsed.era_start_date.cast("date").name("era_start_date"), + (collapsed._max_padded_end - ibis.interval(days=int(gap_days))).cast("date").name("era_end_date"), + ) + + +def compute_drug_eras( + ctx, + *, + drug_codeset_id: int, + gap_days: int, + offset: int, + days_supply_override: int | None, + cohort_person_ids=None, +): + concept_table = ctx.concept_set_table(drug_codeset_id) + + de = ctx.table("drug_exposure") + if cohort_person_ids is not None: + de = de.semi_join( + cohort_person_ids.select(cohort_person_ids.person_id).distinct(), + predicates=[de.person_id == cohort_person_ids.person_id], + ) + + has_source = "drug_source_concept_id" in de.columns + filtered = de.semi_join(concept_table, de.drug_concept_id == concept_table.concept_id) + if has_source: + source_matches = de.semi_join(concept_table, de.drug_source_concept_id == concept_table.concept_id) + filtered = filtered.union(source_matches, distinct=True) + + prepared = filtered.select( + filtered.person_id.cast("int64").name("person_id"), + filtered.drug_exposure_start_date.cast("date").name("start_date"), + _compute_exposure_end_date(filtered, days_supply_override=days_supply_override).name("_exposure_end"), + ) + + return _compute_eras(prepared, gap_days=gap_days, offset=offset) + + +def apply_custom_era_strategy(events, strategy, ctx): + payload = strategy.payload + drug_codeset_id = payload["drug_codeset_id"] + gap_days = payload["gap_days"] + offset = payload["offset"] + days_supply_override = payload.get("days_supply_override") + + if drug_codeset_id is None: + with_bounds = attach_observation_bounds(events, ctx) + return _replace_end_date(events, with_bounds, with_bounds.op_end_date) + + cohort_person_ids = events.select(events.person_id).distinct() + + eras = compute_drug_eras( + ctx, + drug_codeset_id=drug_codeset_id, + gap_days=gap_days, + offset=offset, + days_supply_override=days_supply_override, + cohort_person_ids=cohort_person_ids, + ) + + eras_for_join = eras.select( + eras.person_id.name("_era_person_id"), + eras.era_start_date, + eras.era_end_date, + ) + + with_bounds = attach_observation_bounds(events, ctx) + + joined = with_bounds.left_join( + eras_for_join, + predicates=[ + with_bounds.person_id == eras_for_join._era_person_id, + with_bounds[START_DATE] >= eras_for_join.era_start_date, + with_bounds[START_DATE] <= eras_for_join.era_end_date, + ], + ) + + event_window = ibis.window( + group_by=[joined.person_id, joined.event_id], + order_by=[joined.era_end_date.desc()], + ) + ranked = joined.mutate(_rn=ibis.row_number().over(event_window)) + one_per_event = ranked.filter(ranked._rn == 0) + + effective_end = ibis.coalesce( + one_per_event.era_end_date, + one_per_event.op_end_date, + ) + final_end = ibis.least(effective_end, one_per_event.op_end_date) + + return _replace_end_date(events, one_per_event, final_end) diff --git a/circe/execution/engine/end_strategy.py b/circe/execution/engine/end_strategy.py index a099985b..2585b774 100644 --- a/circe/execution/engine/end_strategy.py +++ b/circe/execution/engine/end_strategy.py @@ -3,10 +3,19 @@ import ibis from ..errors import UnsupportedFeatureError -from ..plan.schema import END_DATE, PERSON_ID, START_DATE +from ..plan.schema import END_DATE, OP_END_DATE, OP_START_DATE, PERSON_ID, START_DATE def attach_observation_bounds(events, ctx): + """Attach observation period bounds to events. + + If events already carry op_start_date/op_end_date from the primary events + stage, use those directly (avoiding a re-join that creates duplicates when + overlapping OPs exist). Falls back to a re-join only if the columns are missing. + """ + if OP_START_DATE in events.columns and OP_END_DATE in events.columns: + return events + observation_period = ctx.table("observation_period").select( PERSON_ID, "observation_period_start_date", @@ -20,8 +29,8 @@ def attach_observation_bounds(events, ctx): ) return joined.select( *[joined[c] for c in events.columns], - observation_period.observation_period_start_date.cast("date").name("op_start_date"), - observation_period.observation_period_end_date.cast("date").name("op_end_date"), + observation_period.observation_period_start_date.cast("date").name(OP_START_DATE), + observation_period.observation_period_end_date.cast("date").name(OP_END_DATE), ).distinct() @@ -39,7 +48,7 @@ def _apply_date_offset_strategy(with_bounds, strategy): ) candidate = base_date + ibis.interval(days=offset) - return ibis.least(candidate, with_bounds.op_end_date) + return ibis.least(candidate, with_bounds[OP_END_DATE]) def _replace_end_date(events, with_bounds, new_end_expr): @@ -57,14 +66,16 @@ def apply_end_strategy(events, strategy, ctx): with_bounds = attach_observation_bounds(events, ctx) if strategy is None: - return _replace_end_date(events, with_bounds, with_bounds.op_end_date) + return _replace_end_date(events, with_bounds, with_bounds[OP_END_DATE]) if strategy.kind == "date_offset": end_date_expr = _apply_date_offset_strategy(with_bounds, strategy) return _replace_end_date(events, with_bounds, end_date_expr) if strategy.kind == "custom_era": - raise UnsupportedFeatureError("Ibis executor end-strategy error: custom_era is not supported.") + from .custom_era import apply_custom_era_strategy + + return apply_custom_era_strategy(events, strategy, ctx) # Fallback: preserve default semantics of op_end_date clipping. - return _replace_end_date(events, with_bounds, with_bounds.op_end_date) + return _replace_end_date(events, with_bounds, with_bounds[OP_END_DATE]) diff --git a/circe/execution/engine/group_demographics.py b/circe/execution/engine/group_demographics.py index bc5920aa..7716ca4c 100644 --- a/circe/execution/engine/group_demographics.py +++ b/circe/execution/engine/group_demographics.py @@ -43,7 +43,7 @@ def _apply_numeric_predicate(expr, predicate): ) -def _apply_date_predicate(expr, predicate): +def _apply_date_predicate(date_expr, predicate): op = (predicate.op or "eq").lower() value = predicate.value extent = predicate.extent @@ -52,19 +52,19 @@ def _apply_date_predicate(expr, predicate): return ibis.literal(True) value_expr = ibis.literal(value).cast("date") - date_expr = expr.cast("date") + if op in {"eq", "="}: - return date_expr == value_expr + return date_expr.cast("date") == value_expr if op in {"neq", "!=", "ne"}: - return date_expr != value_expr + return date_expr.cast("date") != value_expr if op in {"gt", ">"}: - return date_expr > value_expr + return date_expr.cast("date") > value_expr if op in {"gte", ">="}: - return date_expr >= value_expr + return date_expr.cast("date") >= value_expr if op in {"lt", "<"}: - return date_expr < value_expr + return date_expr.cast("date") < value_expr if op in {"lte", "<="}: - return date_expr <= value_expr + return date_expr.cast("date") <= value_expr if op in {"bt", "between"}: if extent is None: raise UnsupportedFeatureError( @@ -74,7 +74,7 @@ def _apply_date_predicate(expr, predicate): extent_expr = ibis.literal(extent).cast("date") lower = ibis.least(value_expr, extent_expr) upper = ibis.greatest(value_expr, extent_expr) - return (date_expr >= lower) & (date_expr <= upper) + return (date_expr.cast("date") >= lower) & (date_expr.cast("date") <= upper) raise UnsupportedFeatureError( f"Ibis executor group evaluation error: unsupported demographic date range op {predicate.op!r}." ) @@ -85,13 +85,24 @@ def _demographic_concept_ids( explicit_ids: tuple[int, ...], codeset_id: int | None, ctx: ExecutionContext, -) -> tuple[int, ...]: - all_ids = list(explicit_ids) +) -> Table | None: + """Resolve concept IDs for a demographic filter. + + Returns an ibis Table with a single ``concept_id`` column, or ``None`` + if neither explicit IDs nor a codeset is provided (meaning no filter). + """ + if not explicit_ids and codeset_id is None: + return None + parts: list[Table] = [] + if explicit_ids: + from ..ibis_compat import literal_column_relation + + parts.append(literal_column_relation(explicit_ids, column_name="concept_id", dtype="int64")) if codeset_id is not None: - for concept_id in ctx.concept_ids_for_codeset(codeset_id): - if concept_id not in all_ids: - all_ids.append(concept_id) - return tuple(all_ids) + parts.append(ctx.concept_set_table(codeset_id).select("concept_id").distinct()) + if len(parts) == 1: + return parts[0] + return parts[0].union(parts[1], distinct=True) def demographic_match_keys( @@ -120,24 +131,24 @@ def demographic_match_keys( codeset_id=demographic.gender_codeset_id, ctx=ctx, ) - if gender_ids: - predicates.append(joined.gender_concept_id.isin(gender_ids)) + if gender_ids is not None: + predicates.append(joined.gender_concept_id.isin(gender_ids.concept_id)) race_ids = _demographic_concept_ids( explicit_ids=demographic.race_concept_ids, codeset_id=demographic.race_codeset_id, ctx=ctx, ) - if race_ids: - predicates.append(joined.race_concept_id.isin(race_ids)) + if race_ids is not None: + predicates.append(joined.race_concept_id.isin(race_ids.concept_id)) ethnicity_ids = _demographic_concept_ids( explicit_ids=demographic.ethnicity_concept_ids, codeset_id=demographic.ethnicity_codeset_id, ctx=ctx, ) - if ethnicity_ids: - predicates.append(joined.ethnicity_concept_id.isin(ethnicity_ids)) + if ethnicity_ids is not None: + predicates.append(joined.ethnicity_concept_id.isin(ethnicity_ids.concept_id)) if demographic.occurrence_start_date is not None: predicates.append( diff --git a/circe/execution/engine/group_keys.py b/circe/execution/engine/group_keys.py index 4de323d6..8510227d 100644 --- a/circe/execution/engine/group_keys.py +++ b/circe/execution/engine/group_keys.py @@ -4,11 +4,17 @@ from ..typing import Table +def _binary_union(tables: list[Table]) -> Table: + if len(tables) == 1: + return tables[0] + mid = len(tables) // 2 + left = _binary_union(tables[:mid]) + right = _binary_union(tables[mid:]) + return left.union(right, distinct=False) + + def union_all(tables: list[Table]) -> Table: - current = tables[0] - for table in tables[1:]: - current = current.union(table, distinct=False) - return current + return _binary_union(tables) def event_keys(events: Table) -> Table: diff --git a/circe/execution/engine/group_operators.py b/circe/execution/engine/group_operators.py index 3ed7c8dd..b2d35ff5 100644 --- a/circe/execution/engine/group_operators.py +++ b/circe/execution/engine/group_operators.py @@ -94,24 +94,43 @@ def group_predicate(match_count_expr, mode: str, count: int | None, child_count: ) +_COMPILED_CORRELATED_EVENTS: dict[tuple[int, int], Table] = {} +"""Cache for :func:`_compile_correlated_events` keyed by ``(backend_id, content_hash)``. + +Identical correlated criteria frequently appear across multiple primary event +criteria within a cohort — compiling them once avoids 350+ duplicate ibis +expression tree constructions for large cohorts. +""" + + def _compile_correlated_events( correlated: NormalizedCorrelatedCriteria, *, criterion_index: int, ctx: ExecutionContext, ) -> Table: + """Compile a correlated criterion to an ibis Table expression. + + The compiled events are independent of *criterion_index* (the position + within the enclosing group), so results are cached by content hash + scoped to the current backend connection. + """ + cache_key = (id(ctx.backend), hash(repr(correlated))) + cached = _COMPILED_CORRELATED_EVENTS.get(cache_key) + if cached is not None: + return cached + event_plan = lower_criterion(correlated.criterion, criterion_index=criterion_index) events = compile_event_plan(event_plan, ctx) nested_group = correlated.criterion.correlated_criteria - if nested_group is None or nested_group.is_empty(): - return events + if nested_group is not None and not nested_group.is_empty(): + from .groups import apply_additional_criteria # noqa: PLC0415 - # Correlated criteria can themselves carry nested correlated criteria. - # Re-apply the same group evaluator used for primary/additional criteria. - from .groups import apply_additional_criteria + events = apply_additional_criteria(events, nested_group, ctx) - return apply_additional_criteria(events, nested_group, ctx) + _COMPILED_CORRELATED_EVENTS[cache_key] = events + return events def correlated_match_keys( diff --git a/circe/execution/engine/groups.py b/circe/execution/engine/groups.py index d630df29..a6a14276 100644 --- a/circe/execution/engine/groups.py +++ b/circe/execution/engine/groups.py @@ -4,7 +4,7 @@ from ..ibis.context import ExecutionContext from ..normalize.groups import NormalizedCriteriaGroup -from ..plan.schema import EVENT_ID, PERSON_ID +from ..plan.schema import EVENT_ID, OP_END_DATE, OP_START_DATE, PERSON_ID from ..typing import Table from .group_demographics import demographic_match_keys from .group_keys import event_keys, union_all @@ -22,32 +22,45 @@ def _evaluate_group( if group.is_empty(): return keys - child_results: list[Table] = [] - index_id = 0 - - for correlated in group.criteria: - correlated_matches = correlated_match_keys( - index_events, - correlated, - criterion_index=index_id, - ctx=ctx, + child_matches: list[Table] = [] + + for index_id, correlated in enumerate(group.criteria): + child_matches.append( + correlated_match_keys( + index_events, + correlated, + criterion_index=index_id, + ctx=ctx, + ) ) - child_results.append(correlated_matches.mutate(index_id=ibis.literal(index_id, type="int64"))) - index_id += 1 + + index_id = len(child_matches) for demographic in group.demographics: - demographic_matches = demographic_match_keys(index_events, demographic, ctx) - child_results.append(demographic_matches.mutate(index_id=ibis.literal(index_id, type="int64"))) + child_matches.append(demographic_match_keys(index_events, demographic, ctx)) index_id += 1 for child_group in group.groups: - child_group_matches = _evaluate_group(index_events, child_group, ctx) - child_results.append(child_group_matches.mutate(index_id=ibis.literal(index_id, type="int64"))) + child_matches.append(_evaluate_group(index_events, child_group, ctx)) index_id += 1 - if not child_results: + if not child_matches: return keys + normalized_mode = (group.mode or "ALL").upper() + if normalized_mode == "ANY": + return union_all(child_matches).distinct() + + if normalized_mode == "AT_LEAST" and group.count is not None and int(group.count) <= 1: + return union_all(child_matches).distinct() + + if len(child_matches) == 1 and normalized_mode == "ALL": + return child_matches[0] + + child_results: list[Table] = [] + for index_id, child_match in enumerate(child_matches): + child_results.append(child_match.mutate(index_id=ibis.literal(index_id, type="int64"))) + unioned = union_all(child_results) group_counts = unioned.group_by(unioned.person_id, unioned.event_id).aggregate( matched_children=unioned.index_id.nunique() @@ -65,7 +78,7 @@ def _evaluate_group( counted.matched_children, group.mode, group.count, - index_id, + len(child_matches), ) return counted.filter(predicate).select( counted.person_id.name(PERSON_ID), @@ -81,7 +94,12 @@ def apply_additional_criteria( if group is None or group.is_empty(): return events - index_events = attach_observation_period(events, ctx) + # Use pre-existing OP bounds from primary events if available, + # avoiding a re-join that creates duplicates when overlapping OPs exist. + if OP_START_DATE in events.columns and OP_END_DATE in events.columns: + index_events = events + else: + index_events = attach_observation_period(events, ctx) matched_keys = _evaluate_group(index_events, group, ctx) filtered = events.join( diff --git a/circe/execution/engine/inclusion.py b/circe/execution/engine/inclusion.py index 7e957842..0e89eacd 100644 --- a/circe/execution/engine/inclusion.py +++ b/circe/execution/engine/inclusion.py @@ -1,6 +1,9 @@ from __future__ import annotations +import ibis + from ..normalize.groups import NormalizedInclusionRule +from .group_keys import event_keys, union_all from .groups import apply_additional_criteria @@ -12,7 +15,33 @@ def apply_inclusion_rules( if not inclusion_rules: return events - included = events - for rule in inclusion_rules: - included = apply_additional_criteria(included, rule.expression, ctx) - return included + active_rule_keys = [] + + for rule_index, rule in enumerate(inclusion_rules): + if rule.expression is None or rule.expression.is_empty(): + continue + + matched = apply_additional_criteria(events, rule.expression, ctx) + active_rule_keys.append(event_keys(matched).mutate(rule_id=ibis.literal(rule_index, type="int64"))) + + if not active_rule_keys: + return events + + if len(active_rule_keys) == 1: + matched_keys = active_rule_keys[0].select("person_id", "event_id") + else: + matched_rules = union_all(active_rule_keys) + matched_counts = matched_rules.group_by("person_id", "event_id").aggregate( + matched_rule_count=matched_rules.rule_id.nunique() + ) + matched_keys = matched_counts.filter( + matched_counts.matched_rule_count == len(active_rule_keys) + ).select("person_id", "event_id") + + included = events.join( + matched_keys, + predicates=[ + (events.person_id == matched_keys.person_id) & (events.event_id == matched_keys.event_id) + ], + ) + return included.select(*[included[c] for c in events.columns]) diff --git a/circe/execution/engine/primary.py b/circe/execution/engine/primary.py index 07f0ec41..11ad0609 100644 --- a/circe/execution/engine/primary.py +++ b/circe/execution/engine/primary.py @@ -7,17 +7,27 @@ from ..ibis.context import ExecutionContext from ..normalize.windows import NormalizedObservationWindow from ..plan.cohort import CohortPlan -from ..plan.schema import DOMAIN, EVENT_ID, PERSON_ID, START_DATE +from ..plan.schema import DOMAIN, EVENT_ID, OP_END_DATE, OP_START_DATE, PERSON_ID, START_DATE from ..typing import Table from .groups import apply_additional_criteria from .limits import apply_result_limit def _union_all(tables): - current = tables[0] - for table in tables[1:]: - current = current.union(table, distinct=False) - return current + if not tables: + raise ValueError("_union_all requires at least one table") + + if len(tables) == 1: + return tables[0] + + # Binary-tree merge: recursively halve the list to produce a balanced + # union tree with O(log n) nesting depth instead of O(n). + # Without this, a cohort with 87 primary criteria would produce 86 levels + # of nested UNION ALL, exceeding DuckDB's query compilation limits. + mid = len(tables) // 2 + left = _union_all(tables[:mid]) + right = _union_all(tables[mid:]) + return left.union(right, distinct=False) def _assign_primary_event_ids(events): @@ -44,7 +54,39 @@ def _apply_observation_window( lower = joined.observation_period_start_date + ibis.interval(days=window.prior_days) upper = joined.observation_period_end_date - ibis.interval(days=window.post_days) filtered = joined.filter((joined[START_DATE] >= lower) & (joined[START_DATE] <= upper)) - return filtered.select(*[filtered[c] for c in events.columns]) + # Carry OP bounds through the pipeline (matching Java behavior). + # Drop any pre-existing op_ columns from earlier stages before re-attaching. + base_cols = [c for c in events.columns if c not in (OP_START_DATE, OP_END_DATE)] + return filtered.select( + *[filtered[c] for c in base_cols], + filtered.observation_period_start_date.cast("date").name(OP_START_DATE), + filtered.observation_period_end_date.cast("date").name(OP_END_DATE), + ) + + +def _attach_op_bounds(events, ctx: ExecutionContext): + """Attach observation period bounds to events without applying an observation window filter. + + This mirrors the Java/R behavior where op_start_date and op_end_date are + always present on primary events for use by end strategy and window constraints. + """ + observation_period = ctx.table("observation_period").select( + PERSON_ID, + "observation_period_start_date", + "observation_period_end_date", + ) + joined = events.join( + observation_period, + (events[PERSON_ID] == observation_period[PERSON_ID]) + & (events[START_DATE] >= observation_period.observation_period_start_date.cast("date")) + & (events[START_DATE] <= observation_period.observation_period_end_date.cast("date")), + ) + base_cols = [c for c in events.columns if c not in (OP_START_DATE, OP_END_DATE)] + return joined.select( + *[joined[c] for c in base_cols], + observation_period.observation_period_start_date.cast("date").name(OP_START_DATE), + observation_period.observation_period_end_date.cast("date").name(OP_END_DATE), + ) def build_primary_events(plan: CohortPlan, ctx: ExecutionContext) -> Table: @@ -64,6 +106,10 @@ def build_primary_events(plan: CohortPlan, ctx: ExecutionContext) -> Table: if plan.observation_window is not None: events = _apply_observation_window(events, ctx, plan.observation_window) + else: + # Always attach OP bounds even without an observation window, + # matching Java behavior where op_end_date is always available. + events = _attach_op_bounds(events, ctx) events = apply_result_limit(events, plan.primary_limit_type) return events diff --git a/circe/execution/ibis/codesets.py b/circe/execution/ibis/codesets.py index f0df82e1..3ca3a34a 100644 --- a/circe/execution/ibis/codesets.py +++ b/circe/execution/ibis/codesets.py @@ -1,229 +1,382 @@ from __future__ import annotations -import hashlib -import json +import contextlib from collections.abc import Callable, Mapping from typing import Any +import ibis + from ..errors import CompilationError -from ..normalize.cohort import NormalizedConceptSet, NormalizedConceptSetItem +from ..normalize.cohort import NormalizedConceptSet from ..plan.schema import CONCEPT_ID from ..typing import IbisBackendLike, Table - -_CACHE_TABLE_NAME = "_circe_codeset_cache" - - -def _compute_cache_key(items: tuple[NormalizedConceptSetItem, ...]) -> str: - """Deterministic SHA-256 hash of sorted concept set items.""" - canonical = sorted( - (item.concept_id, item.is_excluded, item.include_descendants, item.include_mapped) for item in items +from .operations import create_table as _create_table_impl + + +def _literal_select(**columns: int) -> Table: + """Return an ibis table expression that selects literal values. + + Builds ``SELECT val1 AS col1, val2 AS col2`` using ``as_table()`` and + ``mutate()`` so no memtable or local-file staging is required. + """ + items = list(columns.items()) + t = ibis.literal(items[0][1], type="int64").name(items[0][0]).as_table() + for name, val in items[1:]: + t = t.mutate(**{name: ibis.literal(val, type="int64")}) + return t + + +def _empty_table(*, columns: tuple[tuple[str, int], ...]) -> Table: + """Return an ibis expression for an empty table with the given column types. + + Produces ``SELECT ... WHERE FALSE`` -- no local files. + """ + if not columns: + raise ValueError("_empty_table requires at least one column") + t = _literal_select(**dict(columns)) + return t.filter(ibis.literal(False)) + + +def _vocabulary_table( + table_name: str, + *, + vocabulary_schema: str | None, + table_getter: Callable[[str, str | None], Table], +) -> Table: + try: + return table_getter(table_name, vocabulary_schema) + except Exception as exc: + raise CompilationError( + f"Ibis executor compilation error: failed to access vocabulary table '{table_name}'." + ) from exc + + +def _descendant_expression( + ancestor_ids: tuple[int, ...], + *, + table_getter: Callable[[str, str | None], Table], + vocabulary_schema: str | None, +) -> Table: + concept = _vocabulary_table("concept", vocabulary_schema=vocabulary_schema, table_getter=table_getter) + concept_ancestor = _vocabulary_table( + "concept_ancestor", vocabulary_schema=vocabulary_schema, table_getter=table_getter + ) + return ( + concept_ancestor.join(concept, concept_ancestor.descendant_concept_id == concept.concept_id) + .filter(concept_ancestor.ancestor_concept_id.isin(ancestor_ids)) + .filter(concept.invalid_reason.isnull()) + .select(concept_ancestor.descendant_concept_id.name(CONCEPT_ID)) + .distinct() ) - payload = json.dumps(canonical, separators=(",", ":")) - return hashlib.sha256(payload.encode("utf-8")).hexdigest() - -def clear_codeset_cache( - backend: IbisBackendLike, - results_schema: str | None, -) -> None: - """Drop the persistent codeset cache table if it exists.""" - from .operations import create_table, table_exists - if not table_exists(backend, table_name=_CACHE_TABLE_NAME, schema=results_schema): - return +def _mapped_expression( + concept_ids: tuple[int, ...], + *, + table_getter: Callable[[str, str | None], Table], + vocabulary_schema: str | None, +) -> Table: + concept_relationship = _vocabulary_table( + "concept_relationship", vocabulary_schema=vocabulary_schema, table_getter=table_getter + ) + return ( + concept_relationship.filter(concept_relationship.concept_id_2.isin(concept_ids)) + .filter(concept_relationship.relationship_id == "Maps to") + .filter(concept_relationship.invalid_reason.isnull()) + .select(concept_relationship.concept_id_1.name(CONCEPT_ID)) + .distinct() + ) - import ibis - empty = ibis.memtable( - {"cache_key": [], "concept_id": []}, - schema={"cache_key": "string", "concept_id": "int64"}, - ) - create_table(backend, table_name=_CACHE_TABLE_NAME, schema=results_schema, obj=empty, overwrite=True) - - -class CachedConceptSetResolver: - """Resolve concept sets to concrete concept IDs using vocabulary tables.""" - - def __init__( - self, - *, - table_getter: Callable[[str, str | None], Table], - vocabulary_schema: str | None, - concept_sets: Mapping[int, NormalizedConceptSet], - backend: IbisBackendLike | None = None, - results_schema: str | None = None, - use_persistent_cache: bool = False, - ) -> None: - self._table_getter = table_getter - self._vocabulary_schema = vocabulary_schema - self._concept_sets = concept_sets - self._cache: dict[int, tuple[int, ...]] = {} - self._backend = backend - self._results_schema = results_schema - self._use_persistent_cache = ( - use_persistent_cache and backend is not None and results_schema is not None +def _build_codeset_expression( + concept_set: NormalizedConceptSet, + *, + table_getter: Callable[[str, str | None], Table], + vocabulary_schema: str | None, +) -> Table: + """Build a lazy ibis expression for a concept set with include/exclude logic. + + Handles descendants, mapped codes, and exclusions via ibis JOINs. + The database engine performs the expansion at execution time. + Never uses ``ibis.memtable`` -- all leaf values use ``as_table().mutate()`` + to avoid local-file staging on Databricks. + + Batches all ancestor lookups within one concept set into a single + ``concept_ancestor`` JOIN (mirrors Java's ``IN (id1, ..., idN)`` pattern) + rather than issuing one JOIN per item. + """ + # Separate items by (is_excluded) and collect IDs for batched lookups. + include_direct: list[int] = [] + include_desc: list[int] = [] + include_mapped: list[int] = [] + exclude_direct: list[int] = [] + exclude_desc: list[int] = [] + exclude_mapped: list[int] = [] + + for item in concept_set.items: + if item.concept_id is None: + continue + cid = int(item.concept_id) + if item.is_excluded: + exclude_direct.append(cid) + if item.include_descendants: + exclude_desc.append(cid) + if item.include_mapped: + exclude_mapped.append(cid) + else: + include_direct.append(cid) + if item.include_descendants: + include_desc.append(cid) + if item.include_mapped: + include_mapped.append(cid) + + # Build include expression parts with batched vocabulary lookups. + include_parts: list[Table] = [] + if include_direct: + for cid in include_direct: + include_parts.append(_literal_select(concept_id=cid)) + if include_desc: + include_parts.append( + _descendant_expression( + tuple(include_desc), table_getter=table_getter, vocabulary_schema=vocabulary_schema + ) ) - self._persistent_cache_initialized: bool = False - - def resolve_codeset(self, codeset_id: int) -> tuple[int, ...]: - normalized_id = int(codeset_id) - if normalized_id in self._cache: - return self._cache[normalized_id] - - concept_set = self._concept_sets.get(normalized_id) - if concept_set is None or not concept_set.items: - return () - - # L2: persistent cache lookup - cache_key: str | None = None - if self._use_persistent_cache: - cache_key = _compute_cache_key(concept_set.items) - persistent_hit = self._read_persistent_cache(cache_key) - if persistent_hit is not None: - self._cache[normalized_id] = persistent_hit - return persistent_hit - - include_ids: set[int] = set() - exclude_ids: set[int] = set() - for item in concept_set.items: - expanded = self._expand_item(item) - if item.is_excluded: - exclude_ids.update(expanded) - else: - include_ids.update(expanded) - - resolved = tuple(sorted(include_ids - exclude_ids)) - self._cache[normalized_id] = resolved - - # L2: persistent cache write - if self._use_persistent_cache and cache_key is not None and resolved: - self._write_persistent_cache(cache_key, resolved) - - return resolved - - def _expand_item(self, item: NormalizedConceptSetItem) -> set[int]: - base_ids: set[int] = {int(item.concept_id)} - if item.include_descendants: - base_ids.update(self._descendant_ids(base_ids)) - - expanded = set(base_ids) - if item.include_mapped: - expanded.update(self._mapped_ids(base_ids)) - return expanded - - def _vocabulary_table(self, table_name: str) -> Table: - try: - return self._table_getter(table_name, self._vocabulary_schema) - except Exception as exc: # pragma: no cover - backend specific error types - raise CompilationError( - f"Ibis executor compilation error: failed to access vocabulary table '{table_name}'." - ) from exc - - def _descendant_ids(self, ancestor_ids: set[int]) -> set[int]: - if not ancestor_ids: - return set() - - concept = self._vocabulary_table("concept") - concept_ancestor = self._vocabulary_table("concept_ancestor") - query = ( - concept_ancestor.join( - concept, - concept_ancestor.descendant_concept_id == concept.concept_id, + if include_mapped: + include_parts.append( + _mapped_expression( + tuple(include_mapped), table_getter=table_getter, vocabulary_schema=vocabulary_schema + ) + ) + + # Build exclude expression parts with batched vocabulary lookups. + exclude_parts: list[Table] = [] + if exclude_direct: + for cid in exclude_direct: + exclude_parts.append(_literal_select(concept_id=cid)) + if exclude_desc: + exclude_parts.append( + _descendant_expression( + tuple(exclude_desc), table_getter=table_getter, vocabulary_schema=vocabulary_schema ) - .filter(concept_ancestor.ancestor_concept_id.isin(tuple(ancestor_ids))) - .filter(concept.invalid_reason.isnull()) - .select(concept_ancestor.descendant_concept_id.name(CONCEPT_ID)) - .distinct() ) - return self._execute_concept_id_query(query) - - def _mapped_ids(self, input_ids: set[int]) -> set[int]: - if not input_ids: - return set() - - concept_relationship = self._vocabulary_table("concept_relationship") - query = ( - concept_relationship.filter(concept_relationship.concept_id_2.isin(tuple(input_ids))) - .filter(concept_relationship.relationship_id == "Maps to") - .filter(concept_relationship.invalid_reason.isnull()) - .select(concept_relationship.concept_id_1.name(CONCEPT_ID)) - .distinct() + if exclude_mapped: + exclude_parts.append( + _mapped_expression( + tuple(exclude_mapped), table_getter=table_getter, vocabulary_schema=vocabulary_schema + ) ) - return self._execute_concept_id_query(query) - def _execute_concept_id_query(self, query: Table) -> set[int]: + if not include_parts: + return _empty_table(columns=(("concept_id", 0),)) + + # Normalize column nullability for union compatibility across backends. + # _literal_select produces nullable int64 while vocabulary table columns + # (e.g. Databricks concept_ancestor.descendant_concept_id) may be non-nullable. + # Strict backends require exact schema match for UNION ALL. + if len(include_parts) > 1: + include_parts = [p.select(p.concept_id.cast("int64").name(CONCEPT_ID)) for p in include_parts] + + result = _union_all_tables(include_parts) + result = result.distinct() + + if exclude_parts: + if len(exclude_parts) > 1: + exclude_parts = [p.select(p.concept_id.cast("int64").name(CONCEPT_ID)) for p in exclude_parts] + exclude_relation = _union_all_tables(exclude_parts).distinct() + marked = exclude_relation.mutate(_cm=ibis.literal(1, type="int64")) + result = result.join(marked, result.concept_id == marked.concept_id, how="left") + result = result.filter(result._cm.isnull()).drop("_cm") + result = result.select(result.concept_id.name(CONCEPT_ID)) + + return result.select(result.concept_id.name(CONCEPT_ID)) + + +def _union_all_tables(tables: list[Table]) -> Table: + """Union multiple single-column ibis tables using binary-tree merge.""" + if not tables: + raise ValueError("_union_all_tables requires at least one table") + if len(tables) == 1: + return tables[0] + mid = len(tables) // 2 + left = _union_all_tables(tables[:mid]) + right = _union_all_tables(tables[mid:]) + return left.union(right, distinct=False) + + +def _needs_vocabulary_expansion(concept_sets: Mapping[int, NormalizedConceptSet]) -> bool: + for cset in concept_sets.values(): + for item in cset.items: + if item.include_descendants or item.include_mapped: + return True + return False + + +def build_batch_codeset_table( + *, + backend: IbisBackendLike, + concept_sets: Mapping[int, NormalizedConceptSet], + batch_table_name: str, + results_schema: str | None = None, + vocabulary_schema: str | None = None, + temporary: bool = False, +) -> Table: + """Build a ``(codeset_id, concept_id)`` table from multiple concept sets.""" + return build_single_codeset_table( + backend=backend, + concept_sets=concept_sets, + batch_table_name=batch_table_name, + results_schema=results_schema, + vocabulary_schema=vocabulary_schema, + ) + + +def _table_getter_from_backend( + backend: IbisBackendLike, + schema: str, +) -> Callable[[str, str | None], Table]: + def _getter(table_name: str, table_schema: str | None) -> Table: try: - rows = query.execute() - except Exception as exc: # pragma: no cover - backend specific error types - raise CompilationError( - "Ibis executor compilation error: failed executing concept-set expansion query." - ) from exc - - values: list[Any] - if hasattr(rows, "columns"): # pandas DataFrame - values = rows[CONCEPT_ID].tolist() if CONCEPT_ID in rows.columns else rows.iloc[:, 0].tolist() - elif isinstance(rows, (list, tuple, set)): - values = list(rows) - else: - values = [rows] + if table_schema is not None: + return backend.table(table_name, database=table_schema) + except TypeError: + pass + return backend.table(table_name) - output: set[int] = set() - for value in values: - if value is None: - continue - output.add(int(value)) - return output + return _getter - # ------------------------------------------------------------------ - # Persistent cache helpers - # ------------------------------------------------------------------ - def _read_persistent_cache(self, cache_key: str) -> tuple[int, ...] | None: - from .operations import read_table, table_exists +def _read_table( + backend: IbisBackendLike, + *, + table_name: str, + schema: str | None, +) -> Table: + try: + if schema is not None: + return backend.table(table_name, database=schema) + except TypeError: + pass + return backend.table(table_name) + + +def _extract_column(result: Any, col_name: str) -> tuple[int, ...]: + if hasattr(result, "columns"): + values = result[col_name].tolist() if col_name in result.columns else result.iloc[:, 0].tolist() + elif isinstance(result, (list, tuple, set)): + values = list(result) + else: + values = [result] if result is not None else [] + return tuple(int(v) for v in values if v is not None) + + +def _drop_table( + backend: IbisBackendLike, + table_name: str, + schema: str | None, +) -> None: + with contextlib.suppress(Exception): + backend.drop_table(table_name, database=schema, force=True) - try: - if not table_exists(self._backend, table_name=_CACHE_TABLE_NAME, schema=self._results_schema): - return None - tbl = read_table(self._backend, table_name=_CACHE_TABLE_NAME, schema=self._results_schema) - rows = tbl.filter(tbl.cache_key == cache_key).select("concept_id").execute() - if hasattr(rows, "columns"): - values = rows["concept_id"].tolist() - elif isinstance(rows, (list, tuple)): - values = list(rows) - else: - return None - if not values: - return None - return tuple(sorted(int(v) for v in values if v is not None)) - except Exception: - return None - - def _write_persistent_cache(self, cache_key: str, concept_ids: tuple[int, ...]) -> None: - import ibis - - from .operations import create_table, insert_relation, table_exists - try: - data = ibis.memtable( - {"cache_key": [cache_key] * len(concept_ids), "concept_id": list(concept_ids)}, - schema={"cache_key": "string", "concept_id": "int64"}, +def build_single_codeset_table( + *, + backend: IbisBackendLike, + concept_sets: Mapping[int, NormalizedConceptSet], + batch_table_name: str, + results_schema: str | None = None, + vocabulary_schema: str | None = None, + session_prefix: str = "", +) -> Table: + """Build a per-cohort codeset table ``(codeset_id, concept_id)``. + + Like the Java ``#Codesets`` table. All work stays in the database -- + no ``ibis.memtable``, no local-file staging, no ``temp`` tables + (which Databricks / Oracle / BigQuery do not support). + """ + name = f"{session_prefix}{batch_table_name}" + + if not concept_sets: + empty = _empty_table(columns=(("codeset_id", 0), ("concept_id", 0))) + _create_table_impl(backend, table_name=name, schema=results_schema, obj=empty, overwrite=True) + return _read_table(backend, table_name=name, schema=results_schema) + + needs_vocab = _needs_vocabulary_expansion(concept_sets) + + if not needs_vocab: + parts: list[Table] = [] + for cid, cset in concept_sets.items(): + for item in cset.items: + if not item.is_excluded and item.concept_id is not None: + parts.append(_literal_select(codeset_id=int(cid), concept_id=int(item.concept_id))) + if not parts: + empty = _empty_table(columns=(("codeset_id", 0), ("concept_id", 0))) + _create_table_impl(backend, table_name=name, schema=results_schema, obj=empty, overwrite=True) + return _read_table(backend, table_name=name, schema=results_schema) + combined = _union_all_tables(parts) + _create_table_impl(backend, table_name=name, schema=results_schema, obj=combined, overwrite=True) + return _read_table(backend, table_name=name, schema=results_schema) + + table_getter = _table_getter_from_backend(backend, vocabulary_schema or "") + + parts = [] + for cid, cset in concept_sets.items(): + if not cset.items: + continue + has_vocab = any(item.include_descendants or item.include_mapped for item in cset.items) + if has_vocab: + expr = _build_codeset_expression( + cset, table_getter=table_getter, vocabulary_schema=vocabulary_schema ) - if not self._persistent_cache_initialized: - if not table_exists(self._backend, table_name=_CACHE_TABLE_NAME, schema=self._results_schema): - create_table( - self._backend, - table_name=_CACHE_TABLE_NAME, - schema=self._results_schema, - obj=data, - ) - self._persistent_cache_initialized = True - return - self._persistent_cache_initialized = True - insert_relation( - data, - backend=self._backend, - target_table=_CACHE_TABLE_NAME, - target_schema=self._results_schema, + labeled = expr.mutate(codeset_id=ibis.literal(int(cid), type="int64")).select( + "codeset_id", CONCEPT_ID ) - except Exception: - pass + parts.append(labeled) + else: + for item in cset.items: + if not item.is_excluded and item.concept_id is not None: + parts.append(_literal_select(codeset_id=int(cid), concept_id=int(item.concept_id))) + + if not parts: + empty = _empty_table(columns=(("codeset_id", 0), ("concept_id", 0))) + _create_table_impl(backend, table_name=name, schema=results_schema, obj=empty, overwrite=True) + return _read_table(backend, table_name=name, schema=results_schema) + + # Normalize column nullability for union compatibility across backends. + if len(parts) > 1: + parts = [ + p.select( + p.codeset_id.cast("int64").name("codeset_id"), p.concept_id.cast("int64").name(CONCEPT_ID) + ) + for p in parts + ] + + combined = _union_all_tables(parts) + _create_table_impl(backend, table_name=name, schema=results_schema, obj=combined, overwrite=True) + return _read_table(backend, table_name=name, schema=results_schema) + + +def drop_codeset_table( + backend: IbisBackendLike, + *, + batch_table_name: str, + results_schema: str | None = None, +) -> None: + _drop_table(backend, batch_table_name, results_schema) + + +def _filter_by_concept_table( + table: Table, + concept_table: Table, + *, + column: str, + exclude: bool = False, +) -> Table: + """Semi-join (include) or anti-join (exclude) *table* against *concept_table*.""" + if not exclude: + joined = table.join(concept_table, table[column] == concept_table.concept_id) + return joined.select(*[joined[c] for c in table.columns]) + else: + marked = concept_table.mutate(_cm=ibis.literal(1, type="int64")) + joined = table.join(marked, table[column] == marked.concept_id, how="left") + filtered = joined.filter(joined._cm.isnull()) + return filtered.select(*[filtered[c] for c in table.columns]) diff --git a/circe/execution/ibis/compile_steps.py b/circe/execution/ibis/compile_steps.py index 0c6ad844..4d9d9081 100644 --- a/circe/execution/ibis/compile_steps.py +++ b/circe/execution/ibis/compile_steps.py @@ -3,6 +3,7 @@ import ibis from ..errors import CompilationError, UnsupportedFeatureError +from ..ibis_compat import literal_column_relation from ..plan.events import ( ApplyDateAdjustment, FilterByCareSite, @@ -102,24 +103,22 @@ def _apply_date_predicate(expr, predicate: DateRangePredicate): raise CompilationError(f"Ibis executor compilation error: unsupported date range op {predicate.op!r}.") -def _resolve_concept_ids( - *, - direct_ids: tuple[int, ...], - codeset_id: int | None, - ctx: ExecutionContext, -) -> tuple[int, ...]: - all_ids = list(direct_ids) - if codeset_id is not None: - for cid in ctx.concept_ids_for_codeset(codeset_id): - if cid not in all_ids: - all_ids.append(cid) - return tuple(all_ids) - - def _select_original_columns(table, joined): return joined.select(*[joined[c] for c in table.columns]) +def _filter_by_concept_table(table, concept_table, *, column, exclude=False): + """Semi-join (include) or anti-join (exclude) *table* against *concept_table*.""" + if not exclude: + joined = table.join(concept_table, table[column] == concept_table.concept_id) + return _select_original_columns(table, joined) + else: + marked = concept_table.mutate(_cm=ibis.literal(1, type="int64")) + joined = table.join(marked, table[column] == marked.concept_id, how="left") + filtered = joined.filter(joined._cm.isnull()) + return _select_original_columns(table, filtered) + + def _filter_visit_concepts(table, ctx: ExecutionContext, *, step: FilterByVisit): visit = ctx.table("visit_occurrence") visit_lookup = visit.select( @@ -134,14 +133,22 @@ def _filter_visit_concepts(table, ctx: ExecutionContext, *, step: FilterByVisit) table[PERSON_ID] == visit_lookup._visit_person_id, ], ) - concept_ids = _resolve_concept_ids( - direct_ids=step.concept_ids, - codeset_id=step.codeset_id, - ctx=ctx, - ) - predicate = joined._visit_concept_id.isin(concept_ids) - filtered = joined.filter(~predicate if step.exclude else predicate) - return _select_original_columns(table, filtered) + + if step.codeset_id is not None: + concept_table = ctx.concept_set_table(step.codeset_id) + elif step.concept_ids: + concept_table = literal_column_relation(step.concept_ids, column_name="concept_id", dtype="int64") + else: + return _select_original_columns(table, joined) + + if not step.exclude: + joined = joined.join(concept_table, joined._visit_concept_id == concept_table.concept_id) + return _select_original_columns(table, joined) + else: + marked = concept_table.mutate(_cm=ibis.literal(1, type="int64")) + joined = joined.join(marked, joined._visit_concept_id == marked.concept_id, how="left") + joined = joined.filter(joined._cm.isnull()) + return _select_original_columns(table, joined) def _filter_provider_specialty( @@ -159,14 +166,22 @@ def _filter_provider_specialty( provider_lookup, predicates=[table[step.provider_id_column] == provider_lookup._provider_id], ) - concept_ids = _resolve_concept_ids( - direct_ids=step.concept_ids, - codeset_id=step.codeset_id, - ctx=ctx, - ) - predicate = joined._specialty_concept_id.isin(concept_ids) - filtered = joined.filter(~predicate if step.exclude else predicate) - return _select_original_columns(table, filtered) + + if step.codeset_id is not None: + concept_table = ctx.concept_set_table(step.codeset_id) + elif step.concept_ids: + concept_table = literal_column_relation(step.concept_ids, column_name="concept_id", dtype="int64") + else: + return _select_original_columns(table, joined) + + if not step.exclude: + joined = joined.join(concept_table, joined._specialty_concept_id == concept_table.concept_id) + return _select_original_columns(table, joined) + else: + marked = concept_table.mutate(_cm=ibis.literal(1, type="int64")) + joined = joined.join(marked, joined._specialty_concept_id == marked.concept_id, how="left") + joined = joined.filter(joined._cm.isnull()) + return _select_original_columns(table, joined) def _filter_care_site(table, ctx: ExecutionContext, *, step: FilterByCareSite): @@ -179,14 +194,22 @@ def _filter_care_site(table, ctx: ExecutionContext, *, step: FilterByCareSite): care_site_lookup, predicates=[table[step.care_site_id_column] == care_site_lookup._care_site_id], ) - concept_ids = _resolve_concept_ids( - direct_ids=step.concept_ids, - codeset_id=step.codeset_id, - ctx=ctx, - ) - predicate = joined._place_of_service_concept_id.isin(concept_ids) - filtered = joined.filter(~predicate if step.exclude else predicate) - return _select_original_columns(table, filtered) + + if step.codeset_id is not None: + concept_table = ctx.concept_set_table(step.codeset_id) + elif step.concept_ids: + concept_table = literal_column_relation(step.concept_ids, column_name="concept_id", dtype="int64") + else: + return _select_original_columns(table, joined) + + if not step.exclude: + joined = joined.join(concept_table, joined._place_of_service_concept_id == concept_table.concept_id) + return _select_original_columns(table, joined) + else: + marked = concept_table.mutate(_cm=ibis.literal(1, type="int64")) + joined = joined.join(marked, joined._place_of_service_concept_id == marked.concept_id, how="left") + joined = joined.filter(joined._cm.isnull()) + return _select_original_columns(table, joined) def _filter_care_site_location_region( @@ -195,9 +218,7 @@ def _filter_care_site_location_region( *, step: FilterByCareSiteLocationRegion, ): - region_ids = ctx.concept_ids_for_codeset(step.codeset_id) - if not region_ids: - return table.limit(0) + concept_table = ctx.concept_set_table(step.codeset_id) location_history = ctx.table("location_history") history_lookup = location_history.select( @@ -233,8 +254,9 @@ def _filter_care_site_location_region( location_lookup, predicates=[joined_history._history_location_id == location_lookup._location_id], ) - filtered = joined.filter(joined._region_concept_id.isin(region_ids)) - return _select_original_columns(table, filtered) + + joined = joined.join(concept_table, joined._region_concept_id == concept_table.concept_id) + return _select_original_columns(table, joined) def apply_step(step, *, table, source, ctx: ExecutionContext): @@ -253,17 +275,14 @@ def apply_step(step, *, table, source, ctx: ExecutionContext): ) if isinstance(step, FilterByCodeset): - concept_ids = ctx.concept_ids_for_codeset(step.codeset_id) - if not concept_ids: - return table if step.exclude else table.limit(0) - predicate = table[step.column].isin(concept_ids) - return table.filter(~predicate if step.exclude else predicate) + concept_table = ctx.concept_set_table(step.codeset_id) + return _filter_by_concept_table(table, concept_table, column=step.column, exclude=step.exclude) if isinstance(step, FilterByConceptSet): if not step.concept_ids: return table if step.exclude else table.limit(0) - predicate = table[step.column].isin(step.concept_ids) - return table.filter(~predicate if step.exclude else predicate) + concept_table = literal_column_relation(step.concept_ids, column_name="concept_id", dtype="int64") + return _filter_by_concept_table(table, concept_table, column=step.column, exclude=step.exclude) if isinstance(step, FilterByVisit): return _filter_visit_concepts(table, ctx, step=step) diff --git a/circe/execution/ibis/context.py b/circe/execution/ibis/context.py index b7b05ce2..1f09cfcf 100644 --- a/circe/execution/ibis/context.py +++ b/circe/execution/ibis/context.py @@ -1,11 +1,12 @@ from __future__ import annotations from collections.abc import Mapping +from typing import Any from .._dataclass import frozen_slots_dataclass +from ..ibis_compat import literal_rows_relation from ..normalize.cohort import NormalizedConceptSet from ..typing import IbisBackendLike, Table -from .codesets import CachedConceptSetResolver def _table_with_schema_fallback( @@ -27,7 +28,7 @@ class ExecutionContext: cdm_schema: str results_schema: str | None vocabulary_schema: str | None - codeset_resolver: CachedConceptSetResolver + codeset_table: Table def table(self, table_name: str) -> Table: return self._table_from_schema(table_name, self.cdm_schema) @@ -41,37 +42,59 @@ def vocabulary_table(self, table_name: str) -> Table: def _table_from_schema(self, table_name: str, schema: str | None) -> Table: return _table_with_schema_fallback(self.backend, table_name, schema) - def concept_ids_for_codeset(self, codeset_id: int) -> tuple[int, ...]: - return self.codeset_resolver.resolve_codeset(codeset_id) + def concept_set_table(self, codeset_id: int) -> Table: + """Return an ibis Table with a single 'concept_id' column for this codeset. + + References the database-resident batch codeset table. No Python memory + is used for concept IDs -- filtering happens via SQL joins at execution time. + """ + return ( + self.codeset_table.filter(self.codeset_table.codeset_id == codeset_id) + .select("concept_id") + .distinct() + ) + + +def _build_codeset_memtable( + concept_sets: Mapping[int, NormalizedConceptSet], +) -> Table: + """Build a simple table for concept sets with known concept IDs. + + Only handles simple includes (no descendant/mapped expansion needed). + This is a fallback for backward-compatible test usage. + Uses ``literal_rows_relation`` to avoid ``ibis.memtable()``. + """ + rows: list[dict[str, Any]] = [] + for cid, cset in concept_sets.items(): + for item in cset.items: + if not item.is_excluded and item.concept_id is not None: + rows.append({"codeset_id": int(cid), "concept_id": int(item.concept_id)}) + return literal_rows_relation(rows, schema={"codeset_id": "int64", "concept_id": "int64"}) def make_execution_context( *, backend: IbisBackendLike, cdm_schema: str, - concept_sets: Mapping[int, NormalizedConceptSet], + codeset_table: Table | None = None, + concept_sets: Mapping[int, NormalizedConceptSet] | None = None, results_schema: str | None = None, vocabulary_schema: str | None = None, - use_persistent_cache: bool = False, ) -> ExecutionContext: - """Construct an executor context from API-level wiring arguments.""" + """Construct an executor context from API-level wiring arguments. + + Provide *codeset_table* (preferred) for a database-resident codeset + table, or *concept_sets* for backward-compatible single-cohort use. + """ vocabulary_schema = vocabulary_schema or cdm_schema - def _table_getter(table_name: str, schema: str | None) -> Table: - return _table_with_schema_fallback(backend, table_name, schema) + if codeset_table is None: + codeset_table = _build_codeset_memtable(concept_sets or {}) - resolver = CachedConceptSetResolver( - table_getter=_table_getter, - vocabulary_schema=vocabulary_schema, - concept_sets=concept_sets, - backend=backend if use_persistent_cache else None, - results_schema=results_schema if use_persistent_cache else None, - use_persistent_cache=use_persistent_cache, - ) return ExecutionContext( backend=backend, cdm_schema=cdm_schema, results_schema=results_schema, vocabulary_schema=vocabulary_schema, - codeset_resolver=resolver, + codeset_table=codeset_table, ) diff --git a/circe/execution/ibis/operations.py b/circe/execution/ibis/operations.py index c8ba111e..698c64cc 100644 --- a/circe/execution/ibis/operations.py +++ b/circe/execution/ibis/operations.py @@ -83,7 +83,7 @@ def cohort_rows_exist( table = read_table(backend, table_name=cohort_table, schema=results_schema) cohort_id_expr = ibis.literal(int(cohort_id), type="int64") matching = table.filter(table.cohort_definition_id.cast("int64") == cohort_id_expr) - return len(matching.limit(1).execute()) > 0 + return matching.limit(1).count().execute() > 0 except Exception as exc: raise ExecutionError( f"Ibis executor write error: failed checking existing rows for cohort_id={cohort_id}." @@ -119,6 +119,58 @@ def delete_cohort_rows( ) from exc +def insert_rows_via_raw_sql( + backend: IbisBackendLike, + *, + table_name: str, + schema: str | None, + columns: list[str], + rows: list[list], +) -> None: + """Insert rows into a backend table using a raw SQL INSERT VALUES statement. + + Avoids ``ibis.memtable()`` so that backends like Databricks (which + restrict staging/volume paths) can write small payloads without + hitting ``staging_allowed_local_path`` constraints. + """ + from datetime import datetime + + raw_sql = getattr(backend, "raw_sql", None) + if not callable(raw_sql): + raise ExecutionError("Ibis executor write error: backend does not support raw_sql for raw inserts.") + + catalog, database = _catalog_db_tuple(backend, schema) + quoted = getattr(getattr(backend, "compiler", None), "quoted", False) + + table = sg.table(table_name, db=database, catalog=catalog, quoted=quoted) + table_sql = table.sql(dialect=getattr(backend, "name", None) or "duckdb") + + cols_sql = ", ".join( + sg.column(c, quoted=quoted).sql(dialect=getattr(backend, "name", None) or "duckdb") for c in columns + ) + + def _sql_value(v): + if v is None: + return "NULL" + if isinstance(v, (int, float)): + return repr(v) + if isinstance(v, datetime): + return "'" + v.strftime("%Y-%m-%d %H:%M:%S") + "'" + return "'" + str(v).replace("'", "''") + "'" + + values_sql = ", ".join("(" + ", ".join(_sql_value(v) for v in row) + ")" for row in rows) + + statement = f"INSERT INTO {table_sql} ({cols_sql}) VALUES {values_sql}" + + try: + raw_sql(statement) + except Exception as exc: + raise ExecutionError( + "Ibis executor write error: failed inserting rows into " + f"table '{table_name}' in schema '{schema}'." + ) from exc + + def supports_transactional_replace(backend: IbisBackendLike) -> bool: """Return whether cohort-scoped delete+insert can run transactionally.""" return getattr(backend, "name", None) in {"duckdb", "postgres"} diff --git a/circe/execution/ibis/person_filters.py b/circe/execution/ibis/person_filters.py index b46bd998..f8809d23 100644 --- a/circe/execution/ibis/person_filters.py +++ b/circe/execution/ibis/person_filters.py @@ -3,6 +3,7 @@ import ibis from ..errors import CompilationError +from ..ibis_compat import literal_column_relation from ..plan.predicates import NumericRangePredicate from ..plan.schema import PERSON_ID from .context import ExecutionContext @@ -61,18 +62,17 @@ def apply_person_gender_filter( concept_ids: tuple[int, ...], codeset_id: int | None, ): - all_ids = list(concept_ids) if codeset_id is not None: - for cid in ctx.concept_ids_for_codeset(codeset_id): - if cid not in all_ids: - all_ids.append(cid) - - if not all_ids: + concept_table = ctx.concept_set_table(codeset_id) + concept_table = concept_table.select(concept_table.concept_id.name("_pconcept_id")) + elif concept_ids: + concept_table = literal_column_relation(concept_ids, column_name="_pconcept_id", dtype="int64") + else: return table person = ctx.table("person").select(PERSON_ID, "gender_concept_id") joined = table.join(person, table[PERSON_ID] == person[PERSON_ID]) - filtered = joined.filter(joined.gender_concept_id.isin(all_ids)) + filtered = joined.join(concept_table, joined.gender_concept_id == concept_table._pconcept_id) return filtered.select(*[filtered[c] for c in table.columns]) @@ -84,18 +84,17 @@ def _apply_person_concept_filter( concept_ids: tuple[int, ...], codeset_id: int | None, ): - all_ids = list(concept_ids) if codeset_id is not None: - for cid in ctx.concept_ids_for_codeset(codeset_id): - if cid not in all_ids: - all_ids.append(cid) - - if not all_ids: + concept_table = ctx.concept_set_table(codeset_id) + concept_table = concept_table.select(concept_table.concept_id.name("_pconcept_id")) + elif concept_ids: + concept_table = literal_column_relation(concept_ids, column_name="_pconcept_id", dtype="int64") + else: return table person = ctx.table("person").select(PERSON_ID, person_column) joined = table.join(person, table[PERSON_ID] == person[PERSON_ID]) - filtered = joined.filter(joined[person_column].isin(all_ids)) + filtered = joined.join(concept_table, joined[person_column] == concept_table._pconcept_id) return filtered.select(*[filtered[c] for c in table.columns]) diff --git a/circe/execution/lower/common.py b/circe/execution/lower/common.py index 3b474989..989f6da2 100644 --- a/circe/execution/lower/common.py +++ b/circe/execution/lower/common.py @@ -38,6 +38,14 @@ def lower_common_steps(criterion: NormalizedCriterion) -> list[PlanStep]: ) ) + if criterion.source_codeset_id is not None and criterion.source_concept_column is not None: + steps.append( + FilterByCodeset( + column=criterion.source_concept_column, + codeset_id=int(criterion.source_codeset_id), + ) + ) + if criterion.person_filters.gender_concept_ids or criterion.person_filters.gender_codeset_id is not None: steps.append( FilterByPersonGender( diff --git a/circe/execution/normalize/cohort.py b/circe/execution/normalize/cohort.py index b2f657ff..61765b47 100644 --- a/circe/execution/normalize/cohort.py +++ b/circe/execution/normalize/cohort.py @@ -3,7 +3,7 @@ from ...cohortdefinition import CohortExpression from ...vocabulary.concept import ConceptSet from .._dataclass import frozen_slots_dataclass -from ..errors import ExecutionNormalizationError, UnsupportedFeatureError +from ..errors import ExecutionNormalizationError from .collapse import NormalizedCollapseSettings, normalize_collapse_settings from .criteria import NormalizedCriterion, normalize_criterion from .end_strategy import NormalizedEndStrategy, normalize_end_strategy @@ -149,10 +149,6 @@ def normalize_cohort( ) normalized_end_strategy = normalize_end_strategy(expression.end_strategy) - if normalized_end_strategy is not None and normalized_end_strategy.kind == "custom_era": - raise UnsupportedFeatureError( - "Ibis executor normalization error: custom_era end strategy is not supported." - ) return NormalizedCohort( title=expression.title, diff --git a/circe/execution/normalize/criteria.py b/circe/execution/normalize/criteria.py index c3fb6eb3..b7a78282 100644 --- a/circe/execution/normalize/criteria.py +++ b/circe/execution/normalize/criteria.py @@ -62,6 +62,7 @@ class NormalizedCriterion: source_concept_column: str | None visit_occurrence_column: str | None codeset_id: int | None + source_codeset_id: int | None first: bool occurrence_start_date: NormalizedDateRange | None occurrence_end_date: NormalizedDateRange | None @@ -119,6 +120,7 @@ def _build_normalized_criterion( source_concept_column: str | None, visit_occurrence_column: str | None, codeset_id: int | None, + source_codeset_id: int | None = None, first: bool, occurrence_start_date: NormalizedDateRange | None, occurrence_end_date: NormalizedDateRange | None, @@ -135,6 +137,7 @@ def _build_normalized_criterion( source_concept_column=source_concept_column, visit_occurrence_column=visit_occurrence_column, codeset_id=codeset_id, + source_codeset_id=source_codeset_id, first=first, occurrence_start_date=occurrence_start_date, occurrence_end_date=occurrence_end_date, @@ -156,6 +159,7 @@ def _normalize_condition_occurrence(criteria: ConditionOccurrence) -> Normalized source_concept_column="condition_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.condition_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -176,6 +180,7 @@ def _normalize_drug_exposure(criteria: DrugExposure) -> NormalizedCriterion: source_concept_column="drug_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.drug_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -196,6 +201,7 @@ def _normalize_visit_occurrence(criteria: VisitOccurrence) -> NormalizedCriterio source_concept_column="visit_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.visit_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -216,6 +222,7 @@ def _normalize_measurement(criteria: Measurement) -> NormalizedCriterion: source_concept_column="measurement_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.measurement_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -238,6 +245,7 @@ def _normalize_procedure_occurrence( source_concept_column="procedure_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.procedure_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -258,6 +266,7 @@ def _normalize_observation(criteria: Observation) -> NormalizedCriterion: source_concept_column="observation_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.observation_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -278,6 +287,7 @@ def _normalize_visit_detail(criteria: VisitDetail) -> NormalizedCriterion: source_concept_column="visit_detail_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.visit_detail_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.visit_detail_start_date), occurrence_end_date=normalize_date_range(criteria.visit_detail_end_date), @@ -298,6 +308,7 @@ def _normalize_device_exposure(criteria: DeviceExposure) -> NormalizedCriterion: source_concept_column="device_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.device_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -318,6 +329,7 @@ def _normalize_specimen(criteria: Specimen) -> NormalizedCriterion: source_concept_column="specimen_source_concept_id", visit_occurrence_column="visit_occurrence_id", codeset_id=criteria.codeset_id, + source_codeset_id=criteria.specimen_source_concept, first=bool(criteria.first), occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=normalize_date_range(criteria.occurrence_end_date), @@ -338,6 +350,7 @@ def _normalize_death(criteria: Death) -> NormalizedCriterion: source_concept_column="cause_source_concept_id", visit_occurrence_column=None, codeset_id=criteria.codeset_id, + source_codeset_id=criteria.death_source_concept, first=False, occurrence_start_date=normalize_date_range(criteria.occurrence_start_date), occurrence_end_date=None, diff --git a/circe/execution/normalize/end_strategy.py b/circe/execution/normalize/end_strategy.py index 62ff666b..8e034091 100644 --- a/circe/execution/normalize/end_strategy.py +++ b/circe/execution/normalize/end_strategy.py @@ -32,6 +32,7 @@ def normalize_end_strategy( "drug_codeset_id": value.drug_codeset_id, "offset": int(value.offset), "gap_days": int(value.gap_days), + "days_supply_override": value.days_supply_override, }, ) return NormalizedEndStrategy(kind="end_strategy", payload={}) diff --git a/circe/execution/plan/events.py b/circe/execution/plan/events.py index 99652fd7..efe38a46 100644 --- a/circe/execution/plan/events.py +++ b/circe/execution/plan/events.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Union +from typing import Any from .._dataclass import frozen_slots_dataclass from .predicates import DateRangePredicate, NumericRangePredicate @@ -148,27 +148,27 @@ class StandardizeEventShape: end_with: str = "end_date" -PlanStep = Union[ - FilterByCodeset, - FilterByConceptSet, - FilterByDateRange, - FilterByNumericRange, - FilterByText, - JoinLocationRegion, - FilterByVisit, - FilterByVisitDetail, - FilterByProviderSpecialty, - FilterByCareSite, - FilterByCareSiteLocationRegion, - FilterByPersonAge, - FilterByPersonGender, - FilterByPersonRace, - FilterByPersonEthnicity, - KeepFirstPerPerson, - ApplyDateAdjustment, - RestrictToCorrelatedWindow, - StandardizeEventShape, -] +PlanStep = ( + FilterByCodeset + | FilterByConceptSet + | FilterByDateRange + | FilterByNumericRange + | FilterByText + | JoinLocationRegion + | FilterByVisit + | FilterByVisitDetail + | FilterByProviderSpecialty + | FilterByCareSite + | FilterByCareSiteLocationRegion + | FilterByPersonAge + | FilterByPersonGender + | FilterByPersonRace + | FilterByPersonEthnicity + | KeepFirstPerPerson + | ApplyDateAdjustment + | RestrictToCorrelatedWindow + | StandardizeEventShape +) @frozen_slots_dataclass diff --git a/circe/execution/plan/schema.py b/circe/execution/plan/schema.py index 061815f0..375e6fc2 100644 --- a/circe/execution/plan/schema.py +++ b/circe/execution/plan/schema.py @@ -22,6 +22,8 @@ CRITERION_INDEX = "criterion_index" CRITERION_TYPE = "criterion_type" SOURCE_TABLE = "source_table" +OP_START_DATE = "op_start_date" +OP_END_DATE = "op_end_date" STANDARD_EVENT_COLUMNS = ( PERSON_ID, diff --git a/circe/execution/typing.py b/circe/execution/typing.py index edf6ba28..be632a48 100644 --- a/circe/execution/typing.py +++ b/circe/execution/typing.py @@ -1,8 +1,6 @@ from __future__ import annotations -from typing import Any, Protocol - -from typing_extensions import TypeAlias +from typing import Any, Protocol, TypeAlias # Ibis does not currently ship usable type information for its table expressions. # Treat them as `Any` at the compatibility boundary rather than propagating diff --git a/circe/extensions/__init__.py b/circe/extensions/__init__.py index d67a624c..71da52d2 100644 --- a/circe/extensions/__init__.py +++ b/circe/extensions/__init__.py @@ -25,11 +25,12 @@ class WaveformOccurrenceMarkdownRenderer: ... """ +from collections.abc import Callable from pathlib import Path # Forward references to avoid circular imports # Actual imports happen inside methods or with TYPE_CHECKING -from typing import TYPE_CHECKING, Callable, Optional, Union +from typing import TYPE_CHECKING, Optional if TYPE_CHECKING: from ..cohortdefinition.builders.base import CriteriaSqlBuilder @@ -143,7 +144,7 @@ def get_lowerer(self, criteria_cls: type["Criteria"]) -> Optional["LowerFn"]: """ return self._lowerers.get(criteria_cls) - def get_normalizer(self, criteria_cls: type["Criteria"]) -> Optional[NormalizerFn]: + def get_normalizer(self, criteria_cls: type["Criteria"]) -> NormalizerFn | None: """Get the normalizer function for a criteria type. Args: @@ -154,7 +155,7 @@ def get_normalizer(self, criteria_cls: type["Criteria"]) -> Optional[NormalizerF """ return self._normalizers.get(criteria_cls) - def get_template(self, criteria: "Criteria") -> Optional[str]: + def get_template(self, criteria: "Criteria") -> str | None: """Get the markdown template name for a criteria instance. Args: @@ -165,7 +166,7 @@ def get_template(self, criteria: "Criteria") -> Optional[str]: """ return self._markdown_templates.get(type(criteria)) - def get_criteria_class(self, name: str) -> Optional[type["Criteria"]]: + def get_criteria_class(self, name: str) -> type["Criteria"] | None: """Get a registered criteria class by name. Args: @@ -301,7 +302,7 @@ def decorator(cls: type) -> type: return decorator -def template_path(path: Union[str, Path]) -> None: +def template_path(path: str | Path) -> None: """Register a directory as a template search path. This is a convenience function (not a decorator) that adds *path* to the diff --git a/circe/extensions/waveform/criteria.py b/circe/extensions/waveform/criteria.py index d79f9e59..b2d17c95 100644 --- a/circe/extensions/waveform/criteria.py +++ b/circe/extensions/waveform/criteria.py @@ -1,5 +1,3 @@ -from typing import Optional - from pydantic import AliasChoices, Field from circe.cohortdefinition.core import DateRange, NumericRange, TextFilter @@ -20,52 +18,52 @@ class WaveformOccurrence(Criteria): """ # Core concept - type of waveform recording - waveform_occurrence_concept_id: Optional[list[Concept]] = Field( + waveform_occurrence_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("WaveformOccurrenceConceptId", "waveformOccurrenceConceptId"), serialization_alias="WaveformOccurrenceConceptId", ) # Temporal bounds - occurrence_start_datetime: Optional[DateRange] = Field( + occurrence_start_datetime: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceStartDatetime", "occurrenceStartDatetime"), serialization_alias="OccurrenceStartDatetime", ) - occurrence_end_datetime: Optional[DateRange] = Field( + occurrence_end_datetime: DateRange | None = Field( default=None, validation_alias=AliasChoices("OccurrenceEndDatetime", "occurrenceEndDatetime"), serialization_alias="OccurrenceEndDatetime", ) # Visit context - visit_occurrence_id: Optional[NumericRange] = Field( + visit_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("VisitOccurrenceId", "visitOccurrenceId"), serialization_alias="VisitOccurrenceId", ) - visit_detail_id: Optional[NumericRange] = Field( + visit_detail_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("VisitDetailId", "visitDetailId"), serialization_alias="VisitDetailId", ) # File metadata - num_of_files: Optional[NumericRange] = Field( + num_of_files: NumericRange | None = Field( default=None, validation_alias=AliasChoices("NumOfFiles", "numOfFiles"), serialization_alias="NumOfFiles", ) # Source identifiers - waveform_occurrence_source_value: Optional[TextFilter] = Field( + waveform_occurrence_source_value: TextFilter | None = Field( default=None, validation_alias=AliasChoices("WaveformOccurrenceSourceValue", "waveformOccurrenceSourceValue"), serialization_alias="WaveformOccurrenceSourceValue", ) # Sequence/chain filtering - preceding_waveform_occurrence_id: Optional[NumericRange] = Field( + preceding_waveform_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("PrecedingWaveformOccurrenceId", "precedingWaveformOccurrenceId"), serialization_alias="PrecedingWaveformOccurrenceId", @@ -84,43 +82,43 @@ class WaveformRegistry(Criteria): """ # Link to parent occurrence - waveform_occurrence_id: Optional[NumericRange] = Field( + waveform_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("WaveformOccurrenceId", "waveformOccurrenceId"), serialization_alias="WaveformOccurrenceId", ) # File temporal bounds - file_start_datetime: Optional[DateRange] = Field( + file_start_datetime: DateRange | None = Field( default=None, validation_alias=AliasChoices("FileStartDatetime", "fileStartDatetime"), serialization_alias="FileStartDatetime", ) - file_end_datetime: Optional[DateRange] = Field( + file_end_datetime: DateRange | None = Field( default=None, validation_alias=AliasChoices("FileEndDatetime", "fileEndDatetime"), serialization_alias="FileEndDatetime", ) # File format - file_extension_concept_id: Optional[list[Concept]] = Field( + file_extension_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("FileExtensionConceptId", "fileExtensionConceptId"), serialization_alias="FileExtensionConceptId", ) - file_extension_source_value: Optional[TextFilter] = Field( + file_extension_source_value: TextFilter | None = Field( default=None, validation_alias=AliasChoices("FileExtensionSourceValue", "fileExtensionSourceValue"), serialization_alias="FileExtensionSourceValue", ) # Visit context (denormalized for easier querying) - visit_occurrence_id: Optional[NumericRange] = Field( + visit_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("VisitOccurrenceId", "visitOccurrenceId"), serialization_alias="VisitOccurrenceId", ) - visit_detail_id: Optional[NumericRange] = Field( + visit_detail_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("VisitDetailId", "visitDetailId"), serialization_alias="VisitDetailId", @@ -140,62 +138,62 @@ class WaveformChannelMetadata(Criteria): """ # Link to registry file - waveform_registry_id: Optional[NumericRange] = Field( + waveform_registry_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("WaveformRegistryId", "waveformRegistryId"), serialization_alias="WaveformRegistryId", ) # Channel identification - channel_concept_id: Optional[list[Concept]] = Field( + channel_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ChannelConceptId", "channelConceptId"), serialization_alias="ChannelConceptId", ) - waveform_channel_source_value: Optional[TextFilter] = Field( + waveform_channel_source_value: TextFilter | None = Field( default=None, validation_alias=AliasChoices("WaveformChannelSourceValue", "waveformChannelSourceValue"), serialization_alias="WaveformChannelSourceValue", ) # Metadata type (e.g., sampling rate, gain, offset) - metadata_concept_id: Optional[list[Concept]] = Field( + metadata_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("MetadataConceptId", "metadataConceptId"), serialization_alias="MetadataConceptId", ) - metadata_source_value: Optional[TextFilter] = Field( + metadata_source_value: TextFilter | None = Field( default=None, validation_alias=AliasChoices("MetadataSourceValue", "metadataSourceValue"), serialization_alias="MetadataSourceValue", ) # Metadata values (at least one must be populated) - value_as_number: Optional[NumericRange] = Field( + value_as_number: NumericRange | None = Field( default=None, validation_alias=AliasChoices("ValueAsNumber", "valueAsNumber"), serialization_alias="ValueAsNumber", ) - value_as_concept_id: Optional[list[Concept]] = Field( + value_as_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ValueAsConceptId", "valueAsConceptId"), serialization_alias="ValueAsConceptId", ) # Units for numeric values - unit_concept_id: Optional[list[Concept]] = Field( + unit_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("UnitConceptId", "unitConceptId"), serialization_alias="UnitConceptId", ) # Device/procedure linkage - device_exposure_id: Optional[NumericRange] = Field( + device_exposure_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("DeviceExposureId", "deviceExposureId"), serialization_alias="DeviceExposureId", ) - procedure_occurrence_id: Optional[NumericRange] = Field( + procedure_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("ProcedureOccurrenceId", "procedureOccurrenceId"), serialization_alias="ProcedureOccurrenceId", @@ -215,79 +213,79 @@ class WaveformFeature(Criteria): """ # Parent links - waveform_occurrence_id: Optional[NumericRange] = Field( + waveform_occurrence_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("WaveformOccurrenceId", "waveformOccurrenceId"), serialization_alias="WaveformOccurrenceId", ) - waveform_registry_id: Optional[NumericRange] = Field( + waveform_registry_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("WaveformRegistryId", "waveformRegistryId"), serialization_alias="WaveformRegistryId", ) - waveform_channel_metadata_id: Optional[NumericRange] = Field( + waveform_channel_metadata_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("WaveformChannelMetadataId", "waveformChannelMetadataId"), serialization_alias="WaveformChannelMetadataId", ) # Feature type (e.g., heart rate, SpO2, QRS detection) - feature_concept_id: Optional[list[Concept]] = Field( + feature_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("FeatureConceptId", "featureConceptId"), serialization_alias="FeatureConceptId", ) # Algorithm used to derive feature - algorithm_concept_id: Optional[list[Concept]] = Field( + algorithm_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("AlgorithmConceptId", "algorithmConceptId"), serialization_alias="AlgorithmConceptId", ) - algorithm_source_value: Optional[TextFilter] = Field( + algorithm_source_value: TextFilter | None = Field( default=None, validation_alias=AliasChoices("AlgorithmSourceValue", "algorithmSourceValue"), serialization_alias="AlgorithmSourceValue", ) # Temporal window for feature - feature_start_timestamp: Optional[DateRange] = Field( + feature_start_timestamp: DateRange | None = Field( default=None, validation_alias=AliasChoices("FeatureStartTimestamp", "featureStartTimestamp"), serialization_alias="FeatureStartTimestamp", ) - feature_end_timestamp: Optional[DateRange] = Field( + feature_end_timestamp: DateRange | None = Field( default=None, validation_alias=AliasChoices("FeatureEndTimestamp", "featureEndTimestamp"), serialization_alias="FeatureEndTimestamp", ) # Feature values (at least one must be populated) - value_as_number: Optional[NumericRange] = Field( + value_as_number: NumericRange | None = Field( default=None, validation_alias=AliasChoices("ValueAsNumber", "valueAsNumber"), serialization_alias="ValueAsNumber", ) - value_as_concept_id: Optional[list[Concept]] = Field( + value_as_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("ValueAsConceptId", "valueAsConceptId"), serialization_alias="ValueAsConceptId", ) # Units for numeric values - unit_concept_id: Optional[list[Concept]] = Field( + unit_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("UnitConceptId", "unitConceptId"), serialization_alias="UnitConceptId", ) # Links to standard OMOP tables - measurement_id: Optional[NumericRange] = Field( + measurement_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("MeasurementId", "measurementId"), serialization_alias="MeasurementId", ) - observation_id: Optional[NumericRange] = Field( + observation_id: NumericRange | None = Field( default=None, validation_alias=AliasChoices("ObservationId", "observationId"), serialization_alias="ObservationId", diff --git a/circe/io.py b/circe/io.py index af2f1515..98695a44 100644 --- a/circe/io.py +++ b/circe/io.py @@ -10,13 +10,13 @@ import json from collections.abc import Mapping from pathlib import Path -from typing import Any, Union +from typing import Any from .api import cohort_expression_from_json, cohort_expression_from_yaml from .cohortdefinition import CohortExpression from .cohortdefinition.yaml_utils import cohort_expression_to_snake_case -ExpressionInput = Union[CohortExpression, Mapping[str, Any], str, Path] +ExpressionInput = CohortExpression | Mapping[str, Any] | str | Path def load_expression(value: ExpressionInput) -> CohortExpression: diff --git a/circe/vocabulary/concept.py b/circe/vocabulary/concept.py index 6b752b65..0d6e0376 100644 --- a/circe/vocabulary/concept.py +++ b/circe/vocabulary/concept.py @@ -9,7 +9,7 @@ """ from datetime import datetime -from typing import Any, Optional +from typing import Any from pydantic import AliasChoices, BaseModel, ConfigDict, Field, field_validator @@ -26,53 +26,53 @@ class Concept(BaseModel): New schema adds: validStartDate, validEndDate, invalidReason with specific formats. """ - concept_id: Optional[int] = Field( + concept_id: int | None = Field( default=None, validation_alias=AliasChoices("ConceptId", "CONCEPT_ID", "conceptId", "ConceptID"), serialization_alias="CONCEPT_ID", ) - concept_name: Optional[str] = Field( + concept_name: str | None = Field( default=None, validation_alias=AliasChoices("ConceptName", "CONCEPT_NAME", "conceptName"), serialization_alias="CONCEPT_NAME", ) - concept_code: Optional[str] = Field( + concept_code: str | None = Field( default=None, validation_alias=AliasChoices("ConceptCode", "CONCEPT_CODE", "conceptCode"), serialization_alias="CONCEPT_CODE", ) - concept_class_id: Optional[str] = Field( + concept_class_id: str | None = Field( default=None, validation_alias=AliasChoices("ConceptClassId", "CONCEPT_CLASS_ID", "conceptClassId"), serialization_alias="CONCEPT_CLASS_ID", ) - standard_concept: Optional[str] = Field( + standard_concept: str | None = Field( default=None, validation_alias=AliasChoices("StandardConcept", "STANDARD_CONCEPT", "standardConcept"), serialization_alias="STANDARD_CONCEPT", ) - invalid_reason: Optional[str] = Field( + invalid_reason: str | None = Field( default=None, validation_alias=AliasChoices("InvalidReason", "INVALID_REASON", "invalidReason"), serialization_alias="INVALID_REASON", ) - domain_id: Optional[str] = Field( + domain_id: str | None = Field( default=None, validation_alias=AliasChoices("DomainId", "DOMAIN_ID", "domainId"), serialization_alias="DOMAIN_ID", ) - vocabulary_id: Optional[str] = Field( + vocabulary_id: str | None = Field( default=None, validation_alias=AliasChoices("VocabularyId", "VOCABULARY_ID", "vocabularyId"), serialization_alias="VOCABULARY_ID", ) # New schema fields - valid_start_date: Optional[str] = Field( + valid_start_date: str | None = Field( default=None, validation_alias=AliasChoices("validStartDate", "valid_start_date"), serialization_alias="validStartDate", ) - valid_end_date: Optional[str] = Field( + valid_end_date: str | None = Field( default=None, validation_alias=AliasChoices("validEndDate", "valid_end_date"), serialization_alias="validEndDate", @@ -82,7 +82,7 @@ class Concept(BaseModel): @field_validator("standard_concept") @classmethod - def validate_standard_concept(cls, v: Optional[str]) -> Optional[str]: + def validate_standard_concept(cls, v: str | None) -> str | None: """Validate standard_concept is 'S', 'C', or null (relaxed for legacy data).""" # Relaxed validation - warn but don't fail on unexpected values return v @@ -121,11 +121,11 @@ class ConceptSetExpression(BaseModel): (they're sometimes only on the items), so we provide defaults. """ - concept: Optional[Concept] = None + concept: Concept | None = None is_excluded: bool = Field(default=False, alias="isExcluded") include_mapped: bool = Field(default=False, alias="includeMapped") include_descendants: bool = Field(default=False, alias="includeDescendants") - items: Optional[list[ConceptExpressionItem]] = None + items: list[ConceptExpressionItem] | None = None model_config = ConfigDict(populate_by_name=True) @@ -145,7 +145,7 @@ class ConceptSet(BaseModel): description="Unique identifier for the concept set", ) - name: Optional[str] = Field( + name: str | None = Field( default=None, min_length=1, max_length=255, @@ -154,7 +154,7 @@ class ConceptSet(BaseModel): description="Human-readable name for the concept set", ) - expression: Optional[ConceptSetExpression] = Field( + expression: ConceptSetExpression | None = Field( default=None, alias="expression", validation_alias=AliasChoices("expression", "EXPRESSION"), @@ -162,62 +162,62 @@ class ConceptSet(BaseModel): ) # Optional fields for both legacy and new schema - description: Optional[str] = Field( + description: str | None = Field( default=None, max_length=4000, description="Optional detailed description of the concept set purpose and contents", ) # New schema fields (all optional for backward compatibility) - version: Optional[str] = Field( + version: str | None = Field( default=None, description="Version identifier for the concept set (semantic versioning)", ) - created_by: Optional[str] = Field( + created_by: str | None = Field( default=None, alias="createdBy", validation_alias=AliasChoices("createdBy", "created_by"), max_length=255, description="Username or identifier of the concept set creator", ) - created_date: Optional[datetime] = Field( + created_date: datetime | None = Field( default=None, alias="createdDate", validation_alias=AliasChoices("createdDate", "created_date"), description="ISO 8601 timestamp of concept set creation", ) - modified_by: Optional[str] = Field( + modified_by: str | None = Field( default=None, alias="modifiedBy", validation_alias=AliasChoices("modifiedBy", "modified_by"), max_length=255, description="Username or identifier of the last modifier", ) - modified_date: Optional[datetime] = Field( + modified_date: datetime | None = Field( default=None, alias="modifiedDate", validation_alias=AliasChoices("modifiedDate", "modified_date"), description="ISO 8601 timestamp of last modification", ) - created_by_tool: Optional[str] = Field( + created_by_tool: str | None = Field( default=None, alias="createdByTool", validation_alias=AliasChoices("createdByTool", "created_by_tool"), max_length=255, description="Name and version of the tool used to create the concept set", ) - modified_by_tool: Optional[str] = Field( + modified_by_tool: str | None = Field( default=None, alias="modifiedByTool", validation_alias=AliasChoices("modifiedByTool", "modified_by_tool"), max_length=255, description="Name and version of the tool used for the last modification", ) - tags: Optional[list[str]] = Field( + tags: list[str] | None = Field( default=None, description="Optional array of tags for categorization", ) - metadata: Optional[dict[str, Any]] = Field( + metadata: dict[str, Any] | None = Field( default=None, description="Optional additional metadata", ) @@ -226,14 +226,14 @@ class ConceptSet(BaseModel): @field_validator("version") @classmethod - def validate_version(cls, v: Optional[str]) -> Optional[str]: + def validate_version(cls, v: str | None) -> str | None: """Validate semantic versioning pattern if provided (relaxed for legacy compatibility).""" # Relaxed - allow any version string for backward compatibility return v @field_validator("tags") @classmethod - def validate_tags(cls, v: Optional[list[str]]) -> Optional[list[str]]: + def validate_tags(cls, v: list[str] | None) -> list[str] | None: """Validate tags if provided.""" if v is not None: for tag in v: diff --git a/pyproject.toml b/pyproject.toml index 4f0b2a09..d46c64b5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,7 +22,6 @@ classifiers = [ "License :: OSI Approved :: Apache Software License", "Operating System :: OS Independent", "Programming Language :: Python :: 3", - "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", @@ -32,12 +31,13 @@ classifiers = [ "Topic :: Software Development :: Libraries :: Python Modules", "Typing :: Typed", ] -requires-python = ">=3.9" +requires-python = ">=3.10" dependencies = [ "pydantic>=2.0.0", "typing-extensions>=4.0.0", "jinja2>=3.1.0", - "PyYAML>=6.0" + "PyYAML>=6.0", + "ibis-framework[duckdb]>=11.0.0", ] [project.optional-dependencies] @@ -59,17 +59,11 @@ docs = [ "sphinx-rtd-theme>=1.0.0", "myst-parser>=0.18.0", ] -ibis = [ - "ibis-framework>=11.0.0; python_version >= '3.9'", -] -ibis-duckdb = [ - "ibis-framework[duckdb]>=11.0.0; python_version >= '3.9'", -] ibis-postgres = [ - "ibis-framework[postgres]>=11.0.0; python_version >= '3.9'", + "ibis-framework[postgres]>=11.0.0; python_version >= '3.10'", ] ibis-databricks = [ - "ibis-framework[databricks]>=11.0.0; python_version >= '3.9'", + "ibis-framework[databricks]>=11.0.0; python_version >= '3.10'", ] waveform = [ "pydantic>=2.0.0", @@ -97,7 +91,7 @@ circe = ["py.typed"] "circe.extensions.waveform" = ["templates/*.j2"] [tool.mypy] -python_version = "3.9" +python_version = "3.10" warn_return_any = true warn_unused_configs = true disallow_untyped_defs = true diff --git a/tests/execution/phenotype_fixtures/33.json b/tests/execution/phenotype_fixtures/33.json new file mode 100644 index 00000000..aff86431 --- /dev/null +++ b/tests/execution/phenotype_fixtures/33.json @@ -0,0 +1,406 @@ +{ + "cdmVersionRange" : ">=5.0.0", + "PrimaryCriteria" : { + "CriteriaList" : [ + { + "ConditionOccurrence" : { + "CodesetId" : 0, + "ConditionTypeExclude" : false + } + } + ], + "ObservationWindow" : { + "PriorDays" : 0, + "PostDays" : 0 + }, + "PrimaryCriteriaLimit" : { + "Type" : "All" + } + }, + "ConceptSets" : [ + { + "id" : 0, + "name" : "Dementia", + "expression" : { + "items" : [ + { + "concept" : { + "CONCEPT_ID" : 37312036, + "CONCEPT_NAME" : "Aggression due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "788861009", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37312035, + "CONCEPT_NAME" : "Agitation due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "788862002", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4041685, + "CONCEPT_NAME" : "Amyotrophic lateral sclerosis with dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "230258005", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37312031, + "CONCEPT_NAME" : "Anxiety due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "788866004", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37312030, + "CONCEPT_NAME" : "Apathetic behaviour due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "788867008", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 35608576, + "CONCEPT_NAME" : "Behavioral and psychological symptoms of dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "10171000132106", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4092747, + "CONCEPT_NAME" : "Cerebral degeneration presenting primarily with dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "279982005", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4182210, + "CONCEPT_NAME" : "Dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "52448006", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37116464, + "CONCEPT_NAME" : "Dementia caused by heavy metal exposure", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "733184002", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37017549, + "CONCEPT_NAME" : "Dementia co-occurrent with human immunodeficiency virus infection", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "713844000", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4244346, + "CONCEPT_NAME" : "Dialysis dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "9345005", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37311665, + "CONCEPT_NAME" : "Disinhibited behaviour due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "789170003", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4043378, + "CONCEPT_NAME" : "Frontotemporal dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "230270009", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 45765480, + "CONCEPT_NAME" : "Frontotemporal dementia with parkinsonism-17", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "702429008", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 45765477, + "CONCEPT_NAME" : "GRN-related frontotemporal dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "702426001", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 377788, + "CONCEPT_NAME" : "General paresis - neurosyphilis", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "51928006", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 372610, + "CONCEPT_NAME" : "Postconcussion syndrome", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "40425004", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37017247, + "CONCEPT_NAME" : "Presenile dementia co-occurrent with human immunodeficiency virus infection", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "713488003", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : true, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37311890, + "CONCEPT_NAME" : "Psychological symptom due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "789011007", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 37312577, + "CONCEPT_NAME" : "Wandering due to dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "789062005", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4059191, + "CONCEPT_NAME" : "H/O: dementia", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "161465002", + "DOMAIN_ID" : "Observation", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Context-dependent" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + } + ] + } + } + ], + "QualifiedLimit" : { + "Type" : "First" + }, + "ExpressionLimit" : { + "Type" : "All" + }, + "InclusionRules" : [], + "EndStrategy" : { + "DateOffset" : { + "DateField" : "EndDate", + "Offset" : 365 + } + }, + "CensoringCriteria" : [], + "CollapseSettings" : { + "CollapseType" : "ERA", + "EraPad" : 0 + }, + "CensorWindow" : {} +} diff --git a/tests/execution/phenotype_fixtures/54.json b/tests/execution/phenotype_fixtures/54.json new file mode 100644 index 00000000..b74720ac --- /dev/null +++ b/tests/execution/phenotype_fixtures/54.json @@ -0,0 +1,163 @@ +{ + "cdmVersionRange" : ">=5.0.0", + "PrimaryCriteria" : { + "CriteriaList" : [ + { + "ConditionOccurrence" : { + "CodesetId" : 4, + "ConditionTypeExclude" : false + } + } + ], + "ObservationWindow" : { + "PriorDays" : 0, + "PostDays" : 0 + }, + "PrimaryCriteriaLimit" : { + "Type" : "All" + } + }, + "ConceptSets" : [ + { + "id" : 4, + "name" : "Febrile seizure and unspecified seizure ", + "expression" : { + "items" : [ + { + "concept" : { + "CONCEPT_ID" : 444413, + "CONCEPT_NAME" : "Febrile convulsion", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "41497008", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 377091, + "CONCEPT_NAME" : "Seizure", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "91175000", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : false, + "includeMapped" : false + }, + { + "concept" : { + "CONCEPT_ID" : 4196708, + "CONCEPT_NAME" : "Seizure related finding", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "313287004", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + } + ] + } + }, + { + "id" : 6, + "name" : "Febrile seizure", + "expression" : { + "items" : [ + { + "concept" : { + "CONCEPT_ID" : 444413, + "CONCEPT_NAME" : "Febrile convulsion", + "STANDARD_CONCEPT" : "S", + "STANDARD_CONCEPT_CAPTION" : "Standard", + "INVALID_REASON" : "V", + "INVALID_REASON_CAPTION" : "Valid", + "CONCEPT_CODE" : "41497008", + "DOMAIN_ID" : "Condition", + "VOCABULARY_ID" : "SNOMED", + "CONCEPT_CLASS_ID" : "Clinical Finding" + }, + "isExcluded" : false, + "includeDescendants" : true, + "includeMapped" : false + } + ] + } + } + ], + "QualifiedLimit" : { + "Type" : "First" + }, + "ExpressionLimit" : { + "Type" : "All" + }, + "InclusionRules" : [ + { + "name" : "Has febrile seizure diagnosis ", + "expression" : { + "Type" : "ALL", + "CriteriaList" : [ + { + "Criteria" : { + "ConditionOccurrence" : { + "CodesetId" : 6, + "ConditionTypeExclude" : false + } + }, + "StartWindow" : { + "Start" : { + "Days" : 0, + "Coeff" : -1 + }, + "End" : { + "Days" : 42, + "Coeff" : 1 + }, + "UseIndexEnd" : false, + "UseEventEnd" : false + }, + "RestrictVisit" : false, + "IgnoreObservationPeriod" : false, + "Occurrence" : { + "Type" : 2, + "Count" : 1, + "IsDistinct" : false + } + } + ], + "DemographicCriteriaList" : [], + "Groups" : [] + } + } + ], + "EndStrategy" : { + "DateOffset" : { + "DateField" : "EndDate", + "Offset" : 14 + } + }, + "CensoringCriteria" : [], + "CollapseSettings" : { + "CollapseType" : "ERA", + "EraPad" : 0 + }, + "CensorWindow" : {} +} diff --git a/tests/execution/test_api_ibis.py b/tests/execution/test_api_ibis.py index ef0a73e4..db55e52e 100644 --- a/tests/execution/test_api_ibis.py +++ b/tests/execution/test_api_ibis.py @@ -26,8 +26,7 @@ VisitDetail, VisitOccurrence, ) -from circe.cohortdefinition.core import CustomEraStrategy, NumericRange -from circe.execution.errors import UnsupportedFeatureError +from circe.cohortdefinition.core import NumericRange from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem @@ -1213,10 +1212,12 @@ def test_build_cohort_location_region_keeps_repeated_location_history_rows(): assert sorted(result.start_date.astype(str).tolist()) == ["2020-01-01", "2020-02-01"] -def test_build_cohort_rejects_unsupported_features(): - expression = CohortExpression( - primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence()]), - end_strategy=CustomEraStrategy(drug_codeset_id=1, gap_days=30, offset=0), - ) - with pytest.raises(UnsupportedFeatureError, match="custom_era"): - _ = build_cohort(expression, backend=object(), cdm_schema="main") +def test_build_cohort_rejects_unsupported_criteria(): + """Unsupported base criteria type is rejected at normalization time.""" + from circe.cohortdefinition.criteria import Criteria as RawCriteria + from circe.execution.errors import UnsupportedCriterionError + + with pytest.raises(UnsupportedCriterionError): + from circe.execution.normalize.criteria import normalize_criterion + + normalize_criterion(RawCriteria()) diff --git a/tests/execution/test_codeset_batching_verification.py b/tests/execution/test_codeset_batching_verification.py new file mode 100644 index 00000000..9ea5516b --- /dev/null +++ b/tests/execution/test_codeset_batching_verification.py @@ -0,0 +1,326 @@ +"""Verify that the batched codeset optimization produces identical results. + +This test exercises multiple items with include_descendants=True within a single +concept set -- the exact pattern that OPT-1 batches into a single +concept_ancestor JOIN instead of N separate JOINs. +""" + +from __future__ import annotations + +import pytest + +from circe.execution.ibis.codesets import _build_codeset_expression, build_single_codeset_table +from circe.execution.normalize.cohort import NormalizedConceptSet, NormalizedConceptSetItem + + +@pytest.fixture +def vocab_conn(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + conn.create_table( + "concept", + obj=ibis.memtable( + { + "concept_id": [10, 11, 12, 20, 21, 22, 30, 31, 32, 40, 41, 50], + "invalid_reason": [None, None, None, None, None, None, None, None, None, None, None, "D"], + } + ), + overwrite=True, + ) + conn.create_table( + "concept_ancestor", + obj=ibis.memtable( + { + "ancestor_concept_id": [10, 10, 20, 20, 30, 30, 40], + "descendant_concept_id": [11, 12, 21, 22, 31, 32, 50], + } + ), + overwrite=True, + ) + conn.create_table( + "concept_relationship", + obj=ibis.memtable( + { + "concept_id_1": [40, 41, 99], + "concept_id_2": [10, 20, 99], + "relationship_id": ["Maps to", "Maps to", "Maps to"], + "invalid_reason": [None, None, "D"], + } + ), + overwrite=True, + ) + return conn + + +def _table_getter(conn): + def getter(name, schema): + return conn.table(name) + + return getter + + +class TestBatchedDescendantExpansion: + """Test that batching multiple include_descendants items gives correct results.""" + + def test_multiple_descendants_batched(self, vocab_conn): + """Three items with include_descendants=True should resolve to all their descendants.""" + concept_set = NormalizedConceptSet( + set_id=1, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=30, is_excluded=False, include_descendants=True, include_mapped=False + ), + ), + ) + result = _build_codeset_expression( + concept_set, table_getter=_table_getter(vocab_conn), vocabulary_schema=None + ) + rows = set(result.execute()["concept_id"].tolist()) + # Direct {10,20,30} + descendants: 10->{11,12}, 20->{21,22}, 30->{31,32} + # Note: 40->50 but 50 is invalid so not included + assert rows == {10, 11, 12, 20, 21, 22, 30, 31, 32} + + def test_descendants_with_direct_exclusion(self, vocab_conn): + """Excluded item (direct, no descendants) removes it from the final set.""" + concept_set = NormalizedConceptSet( + set_id=2, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=11, is_excluded=True, include_descendants=False, include_mapped=False + ), + ), + ) + result = _build_codeset_expression( + concept_set, table_getter=_table_getter(vocab_conn), vocabulary_schema=None + ) + rows = set(result.execute()["concept_id"].tolist()) + # Include: {10, 11, 12, 20, 21, 22}, Exclude: {11} -> {10, 12, 20, 21, 22} + assert rows == {10, 12, 20, 21, 22} + + def test_excluded_with_descendants_batched(self, vocab_conn): + """Excluded items with include_descendants should batch their ancestor lookup too.""" + concept_set = NormalizedConceptSet( + set_id=3, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=30, is_excluded=True, include_descendants=True, include_mapped=False + ), + ), + ) + result = _build_codeset_expression( + concept_set, table_getter=_table_getter(vocab_conn), vocabulary_schema=None + ) + rows = set(result.execute()["concept_id"].tolist()) + # Include: {10, 11, 12, 20, 21, 22}. Exclude: {30, 31, 32}. No overlap. + assert rows == {10, 11, 12, 20, 21, 22} + + def test_mapped_batched(self, vocab_conn): + """Multiple items with include_mapped should batch the relationship lookup.""" + concept_set = NormalizedConceptSet( + set_id=4, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=False, include_mapped=True + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=False, include_mapped=True + ), + ), + ) + result = _build_codeset_expression( + concept_set, table_getter=_table_getter(vocab_conn), vocabulary_schema=None + ) + rows = set(result.execute()["concept_id"].tolist()) + # Direct: {10, 20}. Mapped: concept_relationship where concept_id_2 IN (10,20) + # -> concept_id_1=40 (maps to 10), concept_id_1=41 (maps to 20) + assert rows == {10, 20, 40, 41} + + def test_descendants_and_mapped_combined(self, vocab_conn): + """Items with both include_descendants and include_mapped batch both lookups.""" + concept_set = NormalizedConceptSet( + set_id=5, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=True, include_mapped=True + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=True, include_mapped=True + ), + ), + ) + result = _build_codeset_expression( + concept_set, table_getter=_table_getter(vocab_conn), vocabulary_schema=None + ) + rows = set(result.execute()["concept_id"].tolist()) + # Direct: {10, 20}. Desc: {11, 12, 21, 22}. Mapped: {40, 41} + assert rows == {10, 11, 12, 20, 21, 22, 40, 41} + + def test_build_single_codeset_table_multiple_sets(self, vocab_conn): + """build_single_codeset_table correctly separates concept sets with batched expansion.""" + concept_sets = { + 1: NormalizedConceptSet( + set_id=1, + items=( + NormalizedConceptSetItem( + concept_id=10, is_excluded=False, include_descendants=True, include_mapped=False + ), + NormalizedConceptSetItem( + concept_id=20, is_excluded=False, include_descendants=True, include_mapped=False + ), + ), + ), + 2: NormalizedConceptSet( + set_id=2, + items=( + NormalizedConceptSetItem( + concept_id=30, is_excluded=False, include_descendants=True, include_mapped=False + ), + ), + ), + } + tbl = build_single_codeset_table( + backend=vocab_conn, + concept_sets=concept_sets, + batch_table_name="__test_batch_verify", + ) + df = tbl.execute() + cs1_ids = set(df[df["codeset_id"] == 1]["concept_id"].tolist()) + cs2_ids = set(df[df["codeset_id"] == 2]["concept_id"].tolist()) + assert cs1_ids == {10, 11, 12, 20, 21, 22} + assert cs2_ids == {30, 31, 32} + vocab_conn.drop_table("__test_batch_verify", force=True) + + +class TestBatchedEndToEnd: + """End-to-end cohort build with batched codeset resolution.""" + + def test_cohort_with_descendants_and_exclusion(self): + """Reproduce the existing test_build_cohort_concept_set_resolves_descendants_and_mapped.""" + import datetime + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from circe.cohortdefinition import ( + CohortExpression, + ConditionOccurrence, + PrimaryCriteria, + ) + from circe.execution.api import build_cohort + from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem + + conn = ibis.duckdb.connect() + S = datetime.date(2020, 1, 1) + E = datetime.date(2020, 12, 31) + conn.create_table( + "person", + obj=ibis.memtable( + {"person_id": [1, 2], "year_of_birth": [1980, 1980], "gender_concept_id": [0, 0]} + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2], + "observation_period_id": [1, 2], + "observation_period_start_date": [S, S], + "observation_period_end_date": [E, E], + } + ), + overwrite=True, + ) + conn.create_table( + "concept", + obj=ibis.memtable( + { + "concept_id": [100, 101, 102, 200, 201], + "invalid_reason": [None, None, "D", None, None], + } + ), + overwrite=True, + ) + conn.create_table( + "concept_ancestor", + obj=ibis.memtable({"ancestor_concept_id": [100, 100], "descendant_concept_id": [101, 102]}), + overwrite=True, + ) + conn.create_table( + "concept_relationship", + obj=ibis.memtable( + { + "concept_id_1": [200, 201], + "concept_id_2": [100, 101], + "relationship_id": ["Maps to", "Maps to"], + "invalid_reason": [None, "D"], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 1, 1, 1, 1, 2], + "condition_occurrence_id": [1000, 1001, 1002, 1003, 1004, 1005], + "condition_concept_id": [100, 101, 102, 200, 201, 999], + "condition_start_date": [S, S, S, S, S, S], + "condition_end_date": [S, S, S, S, S, S], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + ConceptSet( + id=1, + expression=ConceptSetExpression( + items=[ + ConceptSetItem( + concept=Concept(conceptId=100), + includeDescendants=True, + includeMapped=True, + ), + ConceptSetItem( + concept=Concept(conceptId=101), + isExcluded=True, + includeMapped=True, + ), + ] + ), + ) + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + ) + + cohort_result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + assert set(cohort_result.person_id) == {1} + # codeset 1: include 100 (desc=True, mapped=True), exclude 101 (mapped=True) + # Include side: direct 100 + desc of 100: {101, 102(invalid)} + mapped of 100: {200} + # -> valid include = {100, 101, 200} + # Exclude side: direct 101 + mapped of 101: {201(invalid_reason='D')} + # -> valid exclude = {101} + # Final: {100, 101, 200} - {101} = {100, 200} + assert set(cohort_result.concept_id) == {100, 200} diff --git a/tests/execution/test_codeset_resolution.py b/tests/execution/test_codeset_resolution.py new file mode 100644 index 00000000..b7545542 --- /dev/null +++ b/tests/execution/test_codeset_resolution.py @@ -0,0 +1,245 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from circe.cohortdefinition import CohortExpression +from circe.execution.api import build_cohort +from circe.execution.ibis.codesets import build_single_codeset_table +from circe.execution.normalize.cohort import NormalizedConceptSet, NormalizedConceptSetItem + +BENCHMARK_OUTPUT = Path(__file__).resolve().parent.parent.parent / "benchmark_output" +JSON_DIR = Path(__file__).resolve().parent / "phenotype_fixtures" + + +def _seed_minimal_cdm(conn, ibis): + import datetime + + S = datetime.date(2000, 1, 1) + tables = { + "person": {"person_id": [999], "year_of_birth": [1900], "gender_concept_id": [0]}, + "observation_period": { + "person_id": [999], + "observation_period_id": [999], + "observation_period_start_date": [S], + "observation_period_end_date": [S], + }, + "condition_occurrence": { + "person_id": [999], + "condition_occurrence_id": [999], + "condition_concept_id": [0], + "condition_start_date": [S], + "condition_end_date": [S], + }, + "procedure_occurrence": { + "person_id": [999], + "procedure_occurrence_id": [999], + "procedure_concept_id": [0], + "procedure_date": [S], + }, + "measurement": { + "person_id": [999], + "measurement_id": [999], + "measurement_concept_id": [0], + "measurement_date": [S], + }, + "observation": { + "person_id": [999], + "observation_id": [999], + "observation_concept_id": [0], + "observation_date": [S], + }, + "drug_exposure": { + "person_id": [999], + "drug_exposure_id": [999], + "drug_concept_id": [0], + "drug_exposure_start_date": [S], + "drug_exposure_end_date": [S], + }, + "death": {"person_id": [999], "death_date": [S]}, + "visit_occurrence": { + "person_id": [999], + "visit_occurrence_id": [999], + "visit_concept_id": [0], + "visit_start_date": [S], + "visit_end_date": [S], + }, + "specimen": { + "person_id": [999], + "specimen_id": [999], + "specimen_concept_id": [0], + "specimen_date": [S], + }, + "device_exposure": { + "person_id": [999], + "device_exposure_id": [999], + "device_concept_id": [0], + "device_exposure_start_date": [S], + "device_exposure_end_date": [S], + }, + "dose_era": { + "person_id": [999], + "dose_era_id": [999], + "drug_concept_id": [0], + "unit_concept_id": [0], + "dose_value": [0.0], + "dose_era_start_date": [S], + "dose_era_end_date": [S], + }, + "payer_plan_period": { + "person_id": [999], + "payer_plan_period_id": [999], + "payer_plan_period_start_date": [S], + "payer_plan_period_end_date": [S], + }, + "visit_detail": { + "person_id": [999], + "visit_detail_id": [999], + "visit_detail_concept_id": [0], + "visit_detail_start_date": [S], + "visit_detail_end_date": [S], + }, + "condition_era": { + "person_id": [999], + "condition_era_id": [999], + "condition_concept_id": [0], + "condition_era_start_date": [S], + "condition_era_end_date": [S], + "condition_occurrence_count": [1], + }, + "drug_era": { + "person_id": [999], + "drug_era_id": [999], + "drug_concept_id": [0], + "drug_era_start_date": [S], + "drug_era_end_date": [S], + "drug_exposure_count": [1], + "gap_days": [0], + }, + "concept": {"concept_id": [0, 999], "invalid_reason": ["X", None]}, + "concept_ancestor": {"ancestor_concept_id": [999], "descendant_concept_id": [999]}, + "concept_relationship": { + "concept_id_1": [999], + "concept_id_2": [999], + "relationship_id": ["X"], + "invalid_reason": ["X"], + }, + } + for name, obj in tables.items(): + conn.create_table(name, obj=ibis.memtable(obj), overwrite=True) + + +def _make_item(cid: int, *, excluded: bool = False) -> NormalizedConceptSetItem: + return NormalizedConceptSetItem( + concept_id=cid, + is_excluded=excluded, + include_descendants=False, + include_mapped=False, + ) + + +# ------------------------------------------------------------------ +# Multi-exclude collision tests +# ------------------------------------------------------------------ + + +def test_multiple_excludes_no_collision(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + + concept_sets = { + 1: NormalizedConceptSet( + set_id=1, + items=( + _make_item(111), + _make_item(222), + _make_item(333, excluded=True), + _make_item(444, excluded=True), + _make_item(555, excluded=True), + ), + ), + } + + tbl = build_single_codeset_table( + backend=conn, + concept_sets=concept_sets, + batch_table_name="__test_exclude_codesets", + ) + rows = tbl.execute() + assert set(rows["concept_id"].tolist()) == {111, 222} + conn.drop_table("__test_exclude_codesets", force=True) + + +def test_single_exclude_works(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + + concept_sets = { + 1: NormalizedConceptSet( + set_id=1, + items=(_make_item(111), _make_item(222), _make_item(333, excluded=True)), + ), + } + + tbl = build_single_codeset_table( + backend=conn, concept_sets=concept_sets, batch_table_name="__test_exclude_codesets2" + ) + rows = tbl.execute() + assert set(rows["concept_id"].tolist()) == {111, 222} + conn.drop_table("__test_exclude_codesets2", force=True) + + +def test_all_excluded_returns_empty(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + + concept_sets = { + 1: NormalizedConceptSet( + set_id=1, + items=( + _make_item(111, excluded=True), + _make_item(222, excluded=True), + _make_item(333, excluded=True), + ), + ), + } + + tbl = build_single_codeset_table( + backend=conn, concept_sets=concept_sets, batch_table_name="__test_exclude_codesets3" + ) + rows = tbl.execute() + assert len(rows) == 0 + conn.drop_table("__test_exclude_codesets3", force=True) + + +# ------------------------------------------------------------------ +# Phenotype cohort regression tests (recursion / collision fixes) +# ------------------------------------------------------------------ + + +@pytest.mark.parametrize("cohort_id", [33, 54]) +def test_phenotype_cohort_with_exclusions_compiles(cohort_id: int): + """Cohorts 33 and 54 previously failed with recursion errors.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + json_path = JSON_DIR / f"{cohort_id}.json" + if not json_path.exists(): + pytest.skip(f"Phenotype JSON not found: {json_path}") + + expression = CohortExpression.model_validate_json(json_path.read_text()) + + conn = ibis.duckdb.connect() + _seed_minimal_cdm(conn, ibis) + + try: + build_cohort(expression, backend=conn, cdm_schema="main", materialize=False) + except Exception as exc: + pytest.fail(f"Cohort {cohort_id} compilation failed: {exc}") diff --git a/tests/execution/test_codesets_persistent_cache.py b/tests/execution/test_codesets_persistent_cache.py index 087acf8a..c6d19b93 100644 --- a/tests/execution/test_codesets_persistent_cache.py +++ b/tests/execution/test_codesets_persistent_cache.py @@ -2,221 +2,46 @@ import pytest -from circe.execution.ibis.codesets import ( - _CACHE_TABLE_NAME, - CachedConceptSetResolver, - _compute_cache_key, - clear_codeset_cache, -) -from circe.execution.ibis.context import make_execution_context +from circe.execution.ibis.codesets import build_batch_codeset_table from circe.execution.normalize.cohort import NormalizedConceptSet, NormalizedConceptSetItem -# ------------------------------------------------------------------ -# _compute_cache_key tests -# ------------------------------------------------------------------ - -def _make_items(*specs: tuple[int, bool, bool, bool]) -> tuple[NormalizedConceptSetItem, ...]: - return tuple( - NormalizedConceptSetItem( - concept_id=s[0], is_excluded=s[1], include_descendants=s[2], include_mapped=s[3] - ) - for s in specs - ) - - -def test_compute_cache_key_deterministic(): - items = _make_items((1, False, True, False), (2, True, False, True)) - assert _compute_cache_key(items) == _compute_cache_key(items) - - -def test_compute_cache_key_order_independent(): - items_a = _make_items((1, False, True, False), (2, True, False, True)) - items_b = _make_items((2, True, False, True), (1, False, True, False)) - assert _compute_cache_key(items_a) == _compute_cache_key(items_b) - - -def test_compute_cache_key_different_items_different_hash(): - items_a = _make_items((1, False, True, False)) - items_b = _make_items((1, False, False, False)) - assert _compute_cache_key(items_a) != _compute_cache_key(items_b) - - -# ------------------------------------------------------------------ -# Persistent cache integration tests using DuckDB -# ------------------------------------------------------------------ - - -@pytest.fixture -def duckdb_backend(): +def test_build_batch_codeset_table_round_trip(): ibis = pytest.importorskip("ibis") - backend = ibis.duckdb.connect() - backend.raw_sql("CREATE SCHEMA results") - return backend - - -def _concept_set_fixture(): - return { + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + conn.create_table( + "concept", + obj=ibis.memtable( + { + "concept_id": [111, 222, 333], + "invalid_reason": ["X", None, None], + }, + schema={"concept_id": "int64", "invalid_reason": "string"}, + ), + overwrite=True, + ) + + concept_sets = { 1: NormalizedConceptSet( set_id=1, items=( NormalizedConceptSetItem( - concept_id=100, - is_excluded=False, - include_descendants=False, - include_mapped=False, + concept_id=111, is_excluded=False, include_descendants=False, include_mapped=False ), ), - ) + ), } - -def test_persistent_cache_write_and_read(duckdb_backend, monkeypatch): - """First resolve writes to persistent cache; second resolver instance reads from it.""" - concept_sets = _concept_set_fixture() - - resolver1 = CachedConceptSetResolver( - table_getter=lambda name, schema: duckdb_backend.table(name, database=schema), - vocabulary_schema=None, + tbl = build_batch_codeset_table( + backend=conn, concept_sets=concept_sets, - backend=duckdb_backend, - results_schema="results", - use_persistent_cache=True, - ) - - # Bypass vocabulary expansion — just return the concept_id directly - monkeypatch.setattr(resolver1, "_expand_item", lambda item: {item.concept_id}) - - result = resolver1.resolve_codeset(1) - assert result == (100,) - - # Verify the cache table was created with data - cache_tbl = duckdb_backend.table(_CACHE_TABLE_NAME, database="results") - rows = cache_tbl.execute() - assert len(rows) == 1 - - # Second resolver — _expand_item should NOT be called (persistent cache hit) - expand_calls = [] - - resolver2 = CachedConceptSetResolver( - table_getter=lambda name, schema: duckdb_backend.table(name, database=schema), - vocabulary_schema=None, - concept_sets=concept_sets, - backend=duckdb_backend, - results_schema="results", - use_persistent_cache=True, - ) - - def _expand_should_not_be_called(item): - expand_calls.append(item.concept_id) - return {item.concept_id} - - monkeypatch.setattr(resolver2, "_expand_item", _expand_should_not_be_called) - - result2 = resolver2.resolve_codeset(1) - assert result2 == (100,) - assert expand_calls == [], "Expected persistent cache hit — _expand_item should not be called" - - -def test_persistent_cache_disabled_by_default(monkeypatch): - """Without use_persistent_cache=True, no persistent ops happen.""" - concept_sets = _concept_set_fixture() - - resolver = CachedConceptSetResolver( - table_getter=lambda name, schema: None, + batch_table_name="__test_codesets", vocabulary_schema=None, - concept_sets=concept_sets, - ) - - monkeypatch.setattr(resolver, "_expand_item", lambda item: {item.concept_id}) - - result = resolver.resolve_codeset(1) - assert result == (100,) - assert not resolver._use_persistent_cache - - -def test_persistent_cache_read_failure_falls_back_silently(duckdb_backend, monkeypatch): - """If cache read raises, expansion still works.""" - concept_sets = _concept_set_fixture() - - resolver = CachedConceptSetResolver( - table_getter=lambda name, schema: duckdb_backend.table(name, database=schema), - vocabulary_schema=None, - concept_sets=concept_sets, - backend=duckdb_backend, - results_schema="results", - use_persistent_cache=True, - ) - - monkeypatch.setattr(resolver, "_expand_item", lambda item: {item.concept_id}) - - # Force _read_persistent_cache to encounter an error internally by making - # table_exists raise. The method catches all exceptions and returns None. - from circe.execution.ibis import operations as ops - - def _broken_table_exists(*args, **kwargs): - raise RuntimeError("simulated db failure") - - monkeypatch.setattr(ops, "table_exists", _broken_table_exists) - - result = resolver.resolve_codeset(1) - assert result == (100,) - - -def test_clear_codeset_cache(duckdb_backend, monkeypatch): - """clear_codeset_cache empties the cache table.""" - concept_sets = _concept_set_fixture() - - resolver = CachedConceptSetResolver( - table_getter=lambda name, schema: duckdb_backend.table(name, database=schema), - vocabulary_schema=None, - concept_sets=concept_sets, - backend=duckdb_backend, - results_schema="results", - use_persistent_cache=True, - ) - monkeypatch.setattr(resolver, "_expand_item", lambda item: {item.concept_id}) - resolver.resolve_codeset(1) - - # Verify rows exist - cache_tbl = duckdb_backend.table(_CACHE_TABLE_NAME, database="results") - assert len(cache_tbl.execute()) > 0 - - # Clear and verify empty - clear_codeset_cache(duckdb_backend, "results") - cache_tbl = duckdb_backend.table(_CACHE_TABLE_NAME, database="results") - assert len(cache_tbl.execute()) == 0 - - -def test_make_execution_context_threads_persistent_cache(): - """make_execution_context passes persistent cache params to resolver.""" - ibis = pytest.importorskip("ibis") - backend = ibis.duckdb.connect() - - ctx = make_execution_context( - backend=backend, - cdm_schema="main", - concept_sets={}, - results_schema="main", - use_persistent_cache=True, - ) - - assert ctx.codeset_resolver._use_persistent_cache is True - assert ctx.codeset_resolver._backend is backend - assert ctx.codeset_resolver._results_schema == "main" - - -def test_make_execution_context_persistent_cache_disabled_without_results_schema(): - """Persistent cache gracefully disabled when results_schema is None.""" - ibis = pytest.importorskip("ibis") - backend = ibis.duckdb.connect() - - ctx = make_execution_context( - backend=backend, - cdm_schema="main", - concept_sets={}, - results_schema=None, - use_persistent_cache=True, ) + rows = tbl.execute() + assert set(rows["codeset_id"]) == {1} + assert set(rows["concept_id"]) == {111} - assert ctx.codeset_resolver._use_persistent_cache is False + conn.drop_table("__test_codesets", force=True) diff --git a/tests/execution/test_compile_steps_helpers.py b/tests/execution/test_compile_steps_helpers.py index 74a3d878..54413fee 100644 --- a/tests/execution/test_compile_steps_helpers.py +++ b/tests/execution/test_compile_steps_helpers.py @@ -11,7 +11,6 @@ from circe.execution.ibis.compile_steps import ( _apply_date_predicate, _apply_numeric_predicate, - _resolve_concept_ids, apply_step, ) from circe.execution.normalize.windows import NormalizedWindow, NormalizedWindowBound @@ -34,8 +33,9 @@ def __init__(self, conn=None, *, codesets: dict[int, tuple[int, ...]] | None = N self.conn = conn self.codesets = codesets or {} - def concept_ids_for_codeset(self, codeset_id: int) -> tuple[int, ...]: - return self.codesets.get(codeset_id, ()) + def concept_set_table(self, codeset_id: int) -> ibis.Table: + ids = self.codesets.get(codeset_id, ()) + return ibis.memtable({"concept_id": list(ids)}, schema={"concept_id": "int64"}) def table(self, name: str): if self.conn is None: @@ -62,6 +62,7 @@ def _events_table(conn): ], VISIT_OCCURRENCE_ID: [100, 101, 200], "concept_id": [1, 2, 3], + "care_site_id": [10, 20, 30], "text_value": ["alpha", "beta", "gamma"], } ), @@ -131,11 +132,6 @@ def test_apply_date_predicate_rejects_invalid_ranges(): _apply_date_predicate(expr, DateRangePredicate(op="weird", value="2020-01-01", extent=None)) -def test_resolve_concept_ids_deduplicates_codeset_ids(): - ctx = _Context(codesets={1: (2, 3, 4)}) - assert _resolve_concept_ids(direct_ids=(1, 2), codeset_id=1, ctx=ctx) == (1, 2, 3, 4) - - def test_apply_step_covers_text_codeset_concept_and_adjustment_paths(): ibis_mod = pytest.importorskip("ibis") _ = pytest.importorskip("duckdb") @@ -218,7 +214,7 @@ def test_apply_step_covers_keep_first_person_filter_and_error_paths(): table = _events_table(conn) conn.create_table( "person", - obj=ibis_mod.memtable( + obj=ibis.memtable( { PERSON_ID: [1, 2], "gender_concept_id": [8507, 8532], @@ -226,6 +222,30 @@ def test_apply_step_covers_keep_first_person_filter_and_error_paths(): ), overwrite=True, ) + # location_history is needed by FilterByCareSiteLocationRegion + conn.create_table( + "location_history", + obj=ibis.memtable( + { + "entity_id": [1], + "location_id": [1], + "domain_id": ["CARE_SITE"], + "start_date": ["2020-01-01"], + "end_date": ["2020-12-31"], + } + ), + overwrite=True, + ) + conn.create_table( + "location", + obj=ibis.memtable( + { + "location_id": [1], + "region_concept_id": [123], + } + ), + overwrite=True, + ) ctx = _Context(conn, codesets={9: ()}) first = apply_step( diff --git a/tests/execution/test_context_wiring.py b/tests/execution/test_context_wiring.py index 1231f121..2105aa79 100644 --- a/tests/execution/test_context_wiring.py +++ b/tests/execution/test_context_wiring.py @@ -1,100 +1,86 @@ from __future__ import annotations -from types import SimpleNamespace +import pytest -from circe.execution.ibis.codesets import CachedConceptSetResolver +from circe.execution.ibis.codesets import build_single_codeset_table from circe.execution.ibis.context import ExecutionContext, make_execution_context from circe.execution.normalize.cohort import NormalizedConceptSet, NormalizedConceptSetItem -class _BackendWithSchemaSupport: - def __init__(self): - self.calls: list[tuple[str, str | None]] = [] - - def table(self, name: str, database: str | None = None): - self.calls.append((name, database)) - return (name, database) - - -class _BackendWithoutSchemaSupport: - def __init__(self): - self.calls: list[tuple[str, str | None]] = [] - - def table(self, name: str, database: str | None = None): - self.calls.append((name, database)) - if database is not None: - raise TypeError("database kwarg not supported") - return (name, None) +def _make_codeset_table(backend): + return build_single_codeset_table( + backend=backend, + concept_sets={}, + batch_table_name="__test_codesets", + ) def test_make_execution_context_uses_cdm_schema_as_vocabulary_fallback(): - backend = _BackendWithSchemaSupport() + ibis_mod = pytest.importorskip("ibis") + conn = ibis_mod.duckdb.connect() + codeset_table = _make_codeset_table(conn) + ctx = make_execution_context( - backend=backend, - cdm_schema="cdm", - concept_sets={}, + backend=conn, + cdm_schema="main", + codeset_table=codeset_table, ) assert isinstance(ctx, ExecutionContext) - assert ctx.vocabulary_schema == "cdm" - assert isinstance(ctx.codeset_resolver, CachedConceptSetResolver) - assert ctx.table("person") == ("person", "cdm") - assert ctx.concept_ids_for_codeset(999) == () + assert ctx.vocabulary_schema == "main" + + conn.drop_table("__test_codesets", force=True) + +def test_make_execution_context_honors_vocabulary_schema_option(): + ibis_mod = pytest.importorskip("ibis") + conn = ibis_mod.duckdb.connect() + codeset_table = _make_codeset_table(conn) -def test_make_execution_context_honors_vocabulary_schema_option_and_backend_fallback(): - backend = _BackendWithoutSchemaSupport() ctx = make_execution_context( - backend=backend, + backend=conn, cdm_schema="cdm", - concept_sets={}, + codeset_table=codeset_table, vocabulary_schema="vocab", ) assert ctx.vocabulary_schema == "vocab" - assert ctx.vocabulary_table("concept") == ("concept", None) - assert backend.calls == [("concept", "vocab"), ("concept", None)] - -def test_codeset_resolver_caches_expanded_results(monkeypatch): - resolver = CachedConceptSetResolver( - table_getter=lambda name, schema: (name, schema), - vocabulary_schema="vocab", - concept_sets={ - 1: NormalizedConceptSet( - set_id=1, - items=( - NormalizedConceptSetItem( - concept_id=123, - is_excluded=False, - include_descendants=False, - include_mapped=False, - ), - ), - ) - }, - ) - calls: list[int] = [] + conn.drop_table("__test_codesets", force=True) - def _expand(item): - calls.append(item.concept_id) - return {item.concept_id} - monkeypatch.setattr(resolver, "_expand_item", _expand) +def test_codeset_table_returns_filtered_view(): + ibis_mod = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") - assert resolver.resolve_codeset(1) == (123,) - assert resolver.resolve_codeset(1) == (123,) - assert calls == [123] + conn = ibis_mod.duckdb.connect() + concept_sets = { + 1: NormalizedConceptSet( + set_id=1, + items=( + NormalizedConceptSetItem( + concept_id=111, is_excluded=False, include_descendants=False, include_mapped=False + ), + ), + ), + } + + codeset_table = build_single_codeset_table( + backend=conn, + concept_sets=concept_sets, + batch_table_name="__test_codesets2", + ) -def test_codeset_resolver_handles_empty_and_non_dataframe_query_results(): - resolver = CachedConceptSetResolver( - table_getter=lambda name, schema: (name, schema), - vocabulary_schema="vocab", - concept_sets={}, + ctx = make_execution_context( + backend=conn, + cdm_schema="main", + codeset_table=codeset_table, ) - assert resolver._descendant_ids(set()) == set() - assert resolver._mapped_ids(set()) == set() - assert resolver._execute_concept_id_query(SimpleNamespace(execute=lambda: [1, None, 2])) == {1, 2} - assert resolver._execute_concept_id_query(SimpleNamespace(execute=lambda: 3)) == {3} + # concept_set_table should filter by codeset_id + filtered = ctx.concept_set_table(1).execute() + assert len(filtered) == 1 + assert list(filtered["concept_id"]) == [111] + + conn.drop_table("__test_codesets2", force=True) diff --git a/tests/execution/test_correctness_bugs.py b/tests/execution/test_correctness_bugs.py new file mode 100644 index 00000000..7fb4788c --- /dev/null +++ b/tests/execution/test_correctness_bugs.py @@ -0,0 +1,862 @@ +"""Tests for known correctness bugs identified in benchmark comparison. + +Bug 2: End strategy re-joins observation_period, creating duplicates when a person + has overlapping observation periods → overcounting after ERA collapse. + +Bug 3: AdditionalCriteria with VisitOccurrence using EndWindow + UseEventEnd causes + undercounting when visit_end_date is NULL (comparison evaluates to NULL → row dropped). + +Bug 5: Severe undercounting in cohorts with nested CorrelatedCriteria inside + PrimaryCriteria or complex multi-rule inclusion logic. +""" + +from __future__ import annotations + +import pytest + +from circe.api import build_cohort +from circe.cohortdefinition import ( + CohortExpression, + ConditionOccurrence, + CorelatedCriteria, + CriteriaGroup, + Occurrence, + PrimaryCriteria, + ProcedureOccurrence, + VisitOccurrence, + Window, + WindowBound, +) +from circe.cohortdefinition.core import CollapseSettings, DateOffsetStrategy, ResultLimit +from circe.cohortdefinition.criteria import InclusionRule +from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem + + +def _make_concept_set(set_id: int, concept_id: int) -> ConceptSet: + return ConceptSet( + id=set_id, + expression=ConceptSetExpression(items=[ConceptSetItem(concept=Concept(conceptId=concept_id))]), + ) + + +# ────────────────────────────────────────────────────────────────────────────── +# Bug 2: Overlapping observation periods cause duplicate rows in end strategy +# ────────────────────────────────────────────────────────────────────────────── + + +class TestBug2OverlappingObservationPeriods: + """When a person has overlapping observation periods and end strategy uses + DateOffset with EndDate, the re-join to observation_period creates duplicate + rows (one per matching OP) with different op_end_date values, resulting in + different capped end dates. After ERA collapse, this produces more eras than + expected. + + In simple cases, ERA collapse merges the duplicates back (same start_date → + always overlap). However, `attach_observation_bounds` still creates structural + duplication that can cause overcounting in complex pipelines or on backends + with different NULL/tie-breaking semantics. + + Expected behavior: Each event should produce exactly one cohort era regardless + of how many observation periods overlap the event's start_date. + """ + + @pytest.fixture + def conn(self): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + conn = ibis.duckdb.connect() + + # Person 1 has TWO overlapping observation periods with DIFFERENT end dates. + # The shorter OP ends BEFORE the event's end_date + offset, so it caps the + # end date to a different value than the longer OP. + conn.create_table( + "person", + obj=ibis.memtable({"person_id": [1], "year_of_birth": [1980], "gender_concept_id": [8507]}), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 1], + "observation_period_id": [10, 11], + # OP 10: 2019-01-01 to 2020-04-15 (shorter — will cap the end date) + # OP 11: 2020-01-01 to 2021-12-31 (longer — won't cap) + "observation_period_start_date": ["2019-01-01", "2020-01-01"], + "observation_period_end_date": ["2020-04-15", "2021-12-31"], + } + ), + overwrite=True, + ) + # Condition event: start=2020-03-01, end=2020-03-25 + # DateOffset(EndDate, 30): target end_date = 2020-03-25 + 30 = 2020-04-24 + # Via OP 10 (end=2020-04-15): LEAST(2020-04-24, 2020-04-15) = 2020-04-15 + # Via OP 11 (end=2021-12-31): LEAST(2020-04-24, 2021-12-31) = 2020-04-24 + # → Two different end dates for the SAME event → duplicate row + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": ["2020-03-01"], + "condition_end_date": ["2020-03-25"], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + return conn + + def test_attach_observation_bounds_skips_rejoin_when_op_columns_present(self, conn): + """When events already carry op_start_date/op_end_date from primary events, + attach_observation_bounds returns them directly without re-joining. + + This is the fix for Bug 2: the primary events stage now always attaches + OP bounds, so end_strategy and correlated criteria use those instead of + re-joining (which would create duplicates for overlapping OPs). + """ + ibis = pytest.importorskip("ibis") + from circe.execution.engine.end_strategy import attach_observation_bounds + from circe.execution.ibis.context import make_execution_context + + ctx = make_execution_context(backend=conn, cdm_schema="main") + + # Build events WITH op_start_date/op_end_date (as they come from build_primary_events) + events = conn.create_table( + "__test_events_with_op", + obj=ibis.memtable( + { + "person_id": [1], + "event_id": [1], + "start_date": ["2020-03-01"], + "end_date": ["2020-03-25"], + "op_start_date": ["2020-01-01"], + "op_end_date": ["2021-12-31"], + } + ), + overwrite=True, + ) + events = events.mutate( + person_id=events.person_id.cast("int64"), + event_id=events.event_id.cast("int64"), + start_date=events.start_date.cast("date"), + end_date=events.end_date.cast("date"), + op_start_date=events.op_start_date.cast("date"), + op_end_date=events.op_end_date.cast("date"), + ) + + with_bounds = attach_observation_bounds(events, ctx) + result = with_bounds.execute() + + # With OP columns already present, no re-join happens → exactly 1 row + assert len(result) == 1, ( + f"attach_observation_bounds produced {len(result)} rows when OP columns " + "were already present. Should return events directly without re-joining." + ) + + def test_fallback_rejoin_produces_duplicates_for_overlapping_ops(self, conn): + """The fallback re-join path (for events WITHOUT OP columns) still creates + duplicates when overlapping OPs exist. This documents the limitation of the + fallback path, which is NOT used in the normal pipeline (build_primary_events + always attaches OP columns). + """ + ibis = pytest.importorskip("ibis") + from circe.execution.engine.end_strategy import attach_observation_bounds + from circe.execution.ibis.context import make_execution_context + + ctx = make_execution_context(backend=conn, cdm_schema="main") + + # Build events WITHOUT op columns (triggering the fallback re-join) + events = conn.create_table( + "__test_events_no_op", + obj=ibis.memtable( + { + "person_id": [1], + "event_id": [1], + "start_date": ["2020-03-01"], + "end_date": ["2020-03-25"], + } + ), + overwrite=True, + ) + events = events.mutate( + person_id=events.person_id.cast("int64"), + event_id=events.event_id.cast("int64"), + start_date=events.start_date.cast("date"), + end_date=events.end_date.cast("date"), + ) + + with_bounds = attach_observation_bounds(events, ctx) + result = with_bounds.execute() + + # Fallback path: re-join produces 2 rows (one per matching OP) + # This is the known limitation that Bug 2 fix avoids by carrying OP from primary events + assert len(result) == 2, ( + f"Expected fallback re-join to produce 2 rows for overlapping OPs, got {len(result)}." + ) + + def test_single_event_produces_single_era(self, conn): + """One event in overlapping OPs should still produce exactly one cohort era. + + This test passes because ERA collapse merges the duplicate rows back + together (they share the same start_date). But the duplication is still + wasteful and can cause issues on backends with different tie-breaking. + """ + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111)], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="All"), + ), + end_strategy=DateOffsetStrategy(offset=30, date_field="EndDate"), + collapse_settings=CollapseSettings(collapse_type="ERA", era_pad=0), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + assert len(result) == 1 + # End date should use the LONGEST OP (not capped by shorter OP) + assert str(result.iloc[0]["end_date"])[:10] == "2020-04-24" + + def test_two_events_in_overlapping_ops_collapse_correctly(self, conn): + """Two close events that should merge into one era must not be split by OP duplication.""" + ibis = pytest.importorskip("ibis") + # Two events: both fall in both OPs, both get duplicated end dates + # Event 1: start=2020-03-01, end=2020-03-25 → target 2020-04-24 (capped to 2020-04-15 via OP10) + # Event 2: start=2020-03-10, end=2020-03-28 → target 2020-04-27 (capped to 2020-04-15 via OP10) + # These events overlap — should collapse to a single era + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 1], + "condition_occurrence_id": [100, 101], + "condition_concept_id": [111, 111], + "condition_start_date": ["2020-03-01", "2020-03-10"], + "condition_end_date": ["2020-03-25", "2020-03-28"], + "visit_occurrence_id": [10, 10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111)], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="All"), + ), + end_strategy=DateOffsetStrategy(offset=30, date_field="EndDate"), + collapse_settings=CollapseSettings(collapse_type="ERA", era_pad=0), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Both events overlap (event 2 starts before event 1's end+30) → single merged era + # Correct: 1 era from 2020-03-01 to 2020-04-27 + # Bug: OP duplication creates extra rows with different end dates, potentially + # splitting what should be a single era into multiple fragments + assert len(result) == 1, ( + f"Expected 1 merged era but got {len(result)}. " + "Overlapping OPs may have created duplicate rows preventing proper ERA merge." + ) + + +# ────────────────────────────────────────────────────────────────────────────── +# Bug 3: AdditionalCriteria with VisitOccurrence + EndWindow + UseEventEnd +# undercounts when visit_end_date is NULL +# ────────────────────────────────────────────────────────────────────────────── + + +class TestBug3VisitEndWindowUndercounting: + """When AdditionalCriteria requires a VisitOccurrence with EndWindow using + UseEventEnd=true, visits with NULL visit_end_date fail the comparison + (NULL >= date evaluates to NULL/false) and the event is excluded. + + The Java/R reference implementation handles this case (likely via COALESCE + or different NULL semantics), keeping these events in the cohort. + + Pattern from affected cohorts (71, 74, 260-263, 881, 898, 965, 967): + - StartWindow: visit starts on or before the index event + - EndWindow with UseEventEnd=true: visit ends on or after the index event + - This checks "the event occurred during a visit" + """ + + @pytest.fixture + def conn(self): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + conn = ibis.duckdb.connect() + + conn.create_table( + "person", + obj=ibis.memtable( + {"person_id": [1, 2], "year_of_birth": [1980, 1975], "gender_concept_id": [8507, 8532]} + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2], + "observation_period_id": [10, 20], + "observation_period_start_date": ["2019-01-01", "2019-01-01"], + "observation_period_end_date": ["2022-12-31", "2022-12-31"], + } + ), + overwrite=True, + ) + # Two condition events — one for each person + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "condition_occurrence_id": [100, 200], + "condition_concept_id": [111, 111], + "condition_start_date": ["2020-06-15", "2020-06-15"], + "condition_end_date": ["2020-06-20", "2020-06-20"], + "visit_occurrence_id": [1000, 2000], + } + ), + overwrite=True, + ) + # Two visits: + # Person 1: visit with a proper end_date (2020-06-20) — should match + # Person 2: visit with NULL end_date (ongoing visit) — should ALSO match + conn.create_table( + "visit_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "visit_occurrence_id": [1000, 2000], + "visit_concept_id": [9201, 9201], # Inpatient + "visit_start_date": ["2020-06-10", "2020-06-10"], + "visit_end_date": ["2020-06-20", None], # Person 2 has NULL end + "visit_source_concept_id": [0, 0], + } + ), + overwrite=True, + ) + return conn + + def test_null_visit_end_date_excludes_from_end_window(self, conn): + """A visit with NULL end_date is correctly excluded by EndWindow with UseEventEnd=true. + + This matches Java/R behavior: CohortExpressionQueryBuilder.java line 586 uses + A.END_DATE directly without COALESCE. When visit_end_date is NULL, the comparison + `A.END_DATE >= ...` evaluates to NULL → row excluded. This is correct behavior + per OMOP CDM (visit_end_date is a required NOT NULL field). + """ + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111), _make_concept_set(2, 9201)], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="All"), + ), + additional_criteria=CriteriaGroup( + type="ANY", + criteria_list=[ + CorelatedCriteria( + criteria=VisitOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(coeff=-1), # unbounded before + end=WindowBound(days=0, coeff=1), # up to index start + use_index_end=False, + use_event_end=False, + ), + end_window=Window( + start=WindowBound(days=0, coeff=-1), # from index start + end=WindowBound(coeff=1), # unbounded after + use_index_end=False, + use_event_end=True, # compare visit END date + ), + occurrence=Occurrence(type=2, count=1), # at least 1 + ) + ], + ), + end_strategy=DateOffsetStrategy(offset=30, date_field="EndDate"), + collapse_settings=CollapseSettings(collapse_type="ERA", era_pad=0), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Only Person 1 should be in the cohort: + # Person 1: visit (Jun 10 - Jun 20) encompasses condition (Jun 15) ✓ + # Person 2: visit (Jun 10 - NULL) → EndWindow check: NULL >= Jun 15 → NULL → excluded + # This matches Java behavior (no COALESCE on A.END_DATE) + assert len(result) == 1, ( + f"Expected 1 cohort entry (only person with valid visit_end_date) but got {len(result)}." + ) + + def test_visit_not_encompassing_event_excluded(self, conn): + """Visits that genuinely don't encompass the event should still be excluded.""" + ibis = pytest.importorskip("ibis") + # Override: Person 2's visit ENDS BEFORE the condition starts + conn.create_table( + "visit_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "visit_occurrence_id": [1000, 2000], + "visit_concept_id": [9201, 9201], + "visit_start_date": ["2020-06-10", "2020-06-01"], + "visit_end_date": ["2020-06-20", "2020-06-10"], # Person 2 visit ends before condition + "visit_source_concept_id": [0, 0], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111), _make_concept_set(2, 9201)], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="All"), + ), + additional_criteria=CriteriaGroup( + type="ANY", + criteria_list=[ + CorelatedCriteria( + criteria=VisitOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(coeff=-1), + end=WindowBound(days=0, coeff=1), + use_index_end=False, + use_event_end=False, + ), + end_window=Window( + start=WindowBound(days=0, coeff=-1), + end=WindowBound(coeff=1), + use_index_end=False, + use_event_end=True, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + end_strategy=DateOffsetStrategy(offset=30, date_field="EndDate"), + collapse_settings=CollapseSettings(collapse_type="ERA", era_pad=0), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Only person 1 should match — person 2's visit ended before the condition + assert len(result) == 1, ( + f"Expected 1 cohort entry (person 1 only) but got {len(result)}. " + "Person 2's visit ends before condition start and should be excluded." + ) + assert int(result.iloc[0]["person_id"]) == 1 + + +# ────────────────────────────────────────────────────────────────────────────── +# Bug 5a: Nested CorrelatedCriteria inside PrimaryCriteria severe undercount +# ────────────────────────────────────────────────────────────────────────────── + + +class TestBug5NestedCorrelatedCriteria: + """Cohorts with CorrelatedCriteria nested inside PrimaryCriteria (e.g., + "ConditionOccurrence where a ProcedureOccurrence exists within X days") + produce severe undercounting (1-2% of expected results). + + Pattern from cohort 402: PrimaryCriteria has a ConditionOccurrence with + an inline CorrelatedCriteria requiring a ProcedureOccurrence nearby. + """ + + @pytest.fixture + def conn(self): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + conn = ibis.duckdb.connect() + + conn.create_table( + "person", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "year_of_birth": [1980, 1975, 1990], + "gender_concept_id": [8507, 8532, 8507], + } + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "observation_period_id": [10, 20, 30], + "observation_period_start_date": ["2018-01-01", "2018-01-01", "2018-01-01"], + "observation_period_end_date": ["2022-12-31", "2022-12-31", "2022-12-31"], + } + ), + overwrite=True, + ) + # Three persons with conditions + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "condition_occurrence_id": [100, 200, 300], + "condition_concept_id": [111, 111, 111], + "condition_start_date": ["2020-06-01", "2020-07-01", "2020-08-01"], + "condition_end_date": ["2020-06-05", "2020-07-05", "2020-08-05"], + "visit_occurrence_id": [1000, 2000, 3000], + } + ), + overwrite=True, + ) + # Only persons 1 and 2 have a procedure within 7 days of their condition + conn.create_table( + "procedure_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "procedure_occurrence_id": [500, 600], + "procedure_concept_id": [222, 222], + "procedure_date": ["2020-06-03", "2020-07-02"], + "procedure_source_concept_id": [0, 0], + "visit_occurrence_id": [1000, 2000], + } + ), + overwrite=True, + ) + # Need visit_occurrence table for schema completeness + conn.create_table( + "visit_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "visit_occurrence_id": [1000, 2000, 3000], + "visit_concept_id": [9201, 9201, 9201], + "visit_start_date": ["2020-06-01", "2020-07-01", "2020-08-01"], + "visit_end_date": ["2020-06-10", "2020-07-10", "2020-08-10"], + "visit_source_concept_id": [0, 0, 0], + } + ), + overwrite=True, + ) + return conn + + def test_primary_criteria_with_nested_correlated_criteria(self, conn): + """PrimaryCriteria with inline CorrelatedCriteria should correctly filter + primary events to those having a matching correlated event. + + Pattern: "Condition X where Procedure Y occurs within ±7 days" + Only persons 1 and 2 have procedures within 7 days of their condition. + """ + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111), _make_concept_set(2, 222)], + primary_criteria=PrimaryCriteria( + criteria_list=[ + ConditionOccurrence( + codeset_id=1, + correlated_criteria=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ProcedureOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(days=7, coeff=-1), # 7 days before + end=WindowBound(days=7, coeff=1), # 7 days after + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), # at least 1 + ) + ], + ), + ) + ], + primary_limit=ResultLimit(type="First"), + ), + qualified_limit=ResultLimit(type="First"), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Persons 1 and 2 have procedures within 7 days of condition → match + # Person 3 has no procedure at all → excluded by correlated criteria + assert len(result) == 2, ( + f"Expected 2 cohort entries (persons 1 & 2) but got {len(result)}. " + "Nested CorrelatedCriteria in PrimaryCriteria may not be filtering correctly." + ) + result_persons = sorted(result["person_id"].tolist()) + assert result_persons == [1, 2] + + def test_nested_correlated_excludes_non_matching(self, conn): + """Person 3 has no procedure near the condition and should be excluded.""" + pytest.importorskip("ibis") + # Make the window very tight so only person 1 matches (procedure on same day) + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 111), _make_concept_set(2, 222)], + primary_criteria=PrimaryCriteria( + criteria_list=[ + ConditionOccurrence( + codeset_id=1, + correlated_criteria=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ProcedureOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(days=0, coeff=-1), # same day only + end=WindowBound(days=0, coeff=1), + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + ) + ], + primary_limit=ResultLimit(type="First"), + ), + qualified_limit=ResultLimit(type="First"), + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Only person 1's procedure (Jun 3) is within 0 days of condition (Jun 1)? No! + # Person 1: condition=Jun 1, procedure=Jun 3 → 2 days apart → outside ±0 window + # Person 2: condition=Jul 1, procedure=Jul 2 → 1 day apart → outside ±0 window + # Neither should match with a 0-day window + assert len(result) == 0, ( + f"Expected 0 entries with 0-day window but got {len(result)}. " + "No procedures occur on the exact same day as the conditions." + ) + + +# ────────────────────────────────────────────────────────────────────────────── +# Bug 5b: Complex multi-rule inclusion logic causes excessive filtering +# ────────────────────────────────────────────────────────────────────────────── + + +class TestBug5ComplexInclusionRules: + """Cohorts with many inclusion rules (7+) where each rule has complex + correlated criteria may over-filter due to incorrect rule intersection logic. + + Pattern from cohort 726: Multiple inclusion rules, each requiring different + correlated criteria. The ALL logic requires ALL rules to pass for an event + to be included. If any rule's join logic is overly restrictive (e.g., inner + join instead of semi-join), events get dropped. + """ + + @pytest.fixture + def conn(self): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + conn = ibis.duckdb.connect() + + conn.create_table( + "person", + obj=ibis.memtable( + {"person_id": [1, 2], "year_of_birth": [1980, 1975], "gender_concept_id": [8507, 8532]} + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2], + "observation_period_id": [10, 20], + "observation_period_start_date": ["2018-01-01", "2018-01-01"], + "observation_period_end_date": ["2023-12-31", "2023-12-31"], + } + ), + overwrite=True, + ) + # Both persons have the primary condition + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "condition_occurrence_id": [100, 200], + "condition_concept_id": [111, 111], + "condition_start_date": ["2020-06-01", "2020-06-01"], + "condition_end_date": ["2020-06-05", "2020-06-05"], + "visit_occurrence_id": [1000, 2000], + } + ), + overwrite=True, + ) + # Both persons have visit occurrences matching each inclusion rule + conn.create_table( + "visit_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 1, 2, 2], + "visit_occurrence_id": [1000, 1001, 2000, 2001], + "visit_concept_id": [9201, 9202, 9201, 9202], # Inpatient, ER + "visit_start_date": ["2020-06-01", "2020-05-01", "2020-06-01", "2020-05-01"], + "visit_end_date": ["2020-06-10", "2020-05-02", "2020-06-10", "2020-05-02"], + "visit_source_concept_id": [0, 0, 0, 0], + } + ), + overwrite=True, + ) + # Both persons have procedures (for a second inclusion rule) + conn.create_table( + "procedure_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "procedure_occurrence_id": [500, 600], + "procedure_concept_id": [333, 333], + "procedure_date": ["2020-06-02", "2020-06-02"], + "procedure_source_concept_id": [0, 0], + "visit_occurrence_id": [1000, 2000], + } + ), + overwrite=True, + ) + return conn + + def test_multiple_inclusion_rules_all_satisfied(self, conn): + """When a person satisfies ALL inclusion rules, they should remain in the cohort. + + Two inclusion rules: + 1. Must have an inpatient visit (9201) within ±30 days + 2. Must have a procedure (333) within ±30 days + Both persons satisfy both rules. + """ + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 9201), + _make_concept_set(3, 333), + ], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="First"), + ), + inclusion_rules=[ + InclusionRule( + name="Has inpatient visit", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=VisitOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(days=30, coeff=-1), + end=WindowBound(days=30, coeff=1), + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + ), + InclusionRule( + name="Has procedure", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ProcedureOccurrence(codeset_id=3), + start_window=Window( + start=WindowBound(days=30, coeff=-1), + end=WindowBound(days=30, coeff=1), + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + ), + ], + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Both persons have inpatient visit AND procedure within 30 days + assert len(result) == 2, ( + f"Expected 2 cohort entries but got {len(result)}. " + "Multiple inclusion rules may be incorrectly intersecting and dropping valid events." + ) + + def test_one_rule_not_satisfied_excludes_person(self, conn): + """Person failing one inclusion rule should be excluded.""" + ibis = pytest.importorskip("ibis") + # Override procedures — only person 1 has one + conn.create_table( + "procedure_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "procedure_occurrence_id": [500], + "procedure_concept_id": [333], + "procedure_date": ["2020-06-02"], + "procedure_source_concept_id": [0], + "visit_occurrence_id": [1000], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 9201), + _make_concept_set(3, 333), + ], + primary_criteria=PrimaryCriteria( + criteria_list=[ConditionOccurrence(codeset_id=1)], + primary_limit=ResultLimit(type="First"), + ), + inclusion_rules=[ + InclusionRule( + name="Has inpatient visit", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=VisitOccurrence(codeset_id=2), + start_window=Window( + start=WindowBound(days=30, coeff=-1), + end=WindowBound(days=30, coeff=1), + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + ), + InclusionRule( + name="Has procedure", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ProcedureOccurrence(codeset_id=3), + start_window=Window( + start=WindowBound(days=30, coeff=-1), + end=WindowBound(days=30, coeff=1), + use_index_end=False, + use_event_end=False, + ), + occurrence=Occurrence(type=2, count=1), + ) + ], + ), + ), + ], + expression_limit=ResultLimit(type="All"), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + # Person 1: has visit + procedure → passes both rules + # Person 2: has visit but NO procedure → fails rule 2 → excluded + assert len(result) == 1, f"Expected 1 cohort entry (person 1 only) but got {len(result)}." + assert int(result.iloc[0]["person_id"]) == 1 diff --git a/tests/execution/test_custom_era.py b/tests/execution/test_custom_era.py new file mode 100644 index 00000000..90e087c9 --- /dev/null +++ b/tests/execution/test_custom_era.py @@ -0,0 +1,715 @@ +from __future__ import annotations + +from datetime import date + +import pytest + +from circe.api import build_cohort +from circe.cohortdefinition import ( + CohortExpression, + ConditionOccurrence, + DrugExposure, + PrimaryCriteria, +) +from circe.cohortdefinition.core import CustomEraStrategy, ResultLimit +from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem + + +def _make_concept_set(set_id: int, concept_id: int) -> ConceptSet: + return ConceptSet( + id=set_id, + expression=ConceptSetExpression(items=[ConceptSetItem(concept=Concept(conceptId=concept_id))]), + ) + + +def _seed_common_tables(conn, ibis): + conn.create_table( + "person", + obj=ibis.memtable( + { + "person_id": [1], + "year_of_birth": [1980], + "gender_concept_id": [8507], + } + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1], + "observation_period_id": [10], + "observation_period_start_date": [date(2019, 1, 1)], + "observation_period_end_date": [date(2021, 12, 31)], + } + ), + overwrite=True, + ) + + +def test_custom_era_merges_drugs_within_gap(): + """Drug exposures within gap_days merge into one era; cohort end_date reflects it.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1], + "drug_exposure_id": [1, 2], + "drug_concept_id": [222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 2, 1)], + "drug_exposure_end_date": [date(2020, 1, 31), date(2020, 3, 3)], + "days_supply": [0, 0], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 1)], + "condition_end_date": [date(2020, 1, 1)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=30, offset=0), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 1 + assert str(result.iloc[0]["start_date"])[:10] == "2020-01-01" + # exp 1: end=2020-01-31, exp 2: end=2020-03-03 + # gap = 1 <= 30 -> merged era: start=2020-01-01, end=2020-03-03 + assert str(result.iloc[0]["end_date"])[:10] == "2020-03-03" + + +def test_custom_era_no_merge_across_large_gap(): + """Drug exposures beyond gap_days form separate eras; cohort uses nearest era.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1], + "drug_exposure_id": [1, 2], + "drug_concept_id": [222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 2, 1)], + "drug_exposure_end_date": [date(2020, 1, 6), date(2020, 3, 3)], + "days_supply": [0, 0], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 1)], + "condition_end_date": [date(2020, 1, 1)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=5, offset=0), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 1 + assert str(result.iloc[0]["start_date"])[:10] == "2020-01-01" + # exp 1: end=2020-01-06, exp 2: end=2020-03-03 + # gap = 26 > 5 -> separate eras + # cohort start 2020-01-01 matches era 1: end 2020-01-06 + assert str(result.iloc[0]["end_date"])[:10] == "2020-01-06" + + +def test_custom_era_offset_applied(): + """Offset days are added to the drug era end_date.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1], + "drug_exposure_id": [1], + "drug_concept_id": [222], + "drug_exposure_start_date": [date(2020, 1, 1)], + "drug_exposure_end_date": [date(2020, 1, 10)], + "days_supply": [0], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 1)], + "condition_end_date": [date(2020, 1, 1)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=30, offset=7), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 1 + assert str(result.iloc[0]["start_date"])[:10] == "2020-01-01" + # drug effective end: 2020-01-10 (end_date override) + # era: start=2020-01-01, end=2020-01-10+7=2020-01-17 + assert str(result.iloc[0]["end_date"])[:10] == "2020-01-17" + + +def test_custom_era_no_matching_drugs(): + """No matching drug exposures -> fall back to observation_period_end_date.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 15)], + "condition_end_date": [date(2020, 1, 15)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [], + "drug_exposure_id": [], + "drug_concept_id": [], + "drug_exposure_start_date": [], + "drug_exposure_end_date": [], + "days_supply": [], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 999), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=30, offset=0), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 1 + assert str(result.iloc[0]["start_date"])[:10] == "2020-01-15" + # No matching drugs -> end_date = observation_period_end_date = 2021-12-31 + assert str(result.iloc[0]["end_date"])[:10] == "2021-12-31" + + +def test_custom_era_with_drug_exposure_as_primary(): + """Custom era works with DrugExposure as the primary criterion.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1], + "drug_exposure_id": [1, 2], + "drug_concept_id": [222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 2, 1)], + "drug_exposure_end_date": [date(2020, 1, 31), date(2020, 3, 3)], + "days_supply": [0, 0], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 222)], + primary_criteria=PrimaryCriteria(criteria_list=[DrugExposure(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=1, gap_days=30, offset=0), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + # With primary_limit_type="all", both drug exposures produce cohort entries. + # Both entries get end_date from the merged drug era (2020-03-03). + assert len(result) == 2 + start_dates = sorted(result["start_date"].astype(str).tolist()) + assert start_dates == ["2020-01-01", "2020-02-01"] + assert all(str(d)[:10] == "2020-03-03" for d in result["end_date"]) + + +def test_compute_drug_eras_matches_java_sql_logic(): + """compute_drug_eras ibis output matches equivalent raw SQL (Java template translated to DuckDB).""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from types import SimpleNamespace + + from circe.execution.engine.custom_era import compute_drug_eras + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + # 5 exposures for person 1, with gap_days=7, offset=3. + # Exposure end_dates are set explicitly so COALESCE is predictable. + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1, 1, 1, 1], + "drug_exposure_id": [1, 2, 3, 4, 5], + "drug_concept_id": [222, 222, 222, 222, 222], + "drug_exposure_start_date": [ + date(2020, 1, 1), + date(2020, 1, 10), + date(2020, 3, 1), + date(2020, 3, 20), + date(2020, 5, 1), + ], + "drug_exposure_end_date": [ + date(2020, 1, 6), + date(2020, 2, 9), + date(2020, 3, 21), + date(2020, 3, 30), + date(2020, 5, 15), + ], + "days_supply": [0, 0, 0, 0, 0], + } + ), + overwrite=True, + ) + + ctx = SimpleNamespace( + table=lambda name: conn.table(name), + concept_set_table=lambda cid: ibis.memtable( + {"concept_id": [222] if cid == 2 else []}, schema={"concept_id": "int64"} + ), + ) + + # --- ibis path --- + ibis_result = compute_drug_eras( + ctx, drug_codeset_id=2, gap_days=7, offset=3, days_supply_override=None + ).execute() + ibis_result = ibis_result.sort_values(["person_id", "era_start_date"]).reset_index(drop=True) + + # --- raw SQL path (Java template core logic, DuckDB dialect) --- + # Java template uses: COALESCE(end, start+days_supply, start+1) + # then pads by (gap_days + offset), groups by cumulative-max-over-preceding, + # and finally subtracts gap_days from max(end) to leave only offset. + gap = 7 + off = 3 + + sql = f""" + WITH exposures AS ( + SELECT + person_id::INTEGER AS person_id, + drug_exposure_start_date::DATE AS start_date, + COALESCE( + drug_exposure_end_date::DATE, + drug_exposure_start_date::DATE + days_supply::INTEGER, + drug_exposure_start_date::DATE + 1 + ) + {gap + off} AS padded_end + FROM drug_exposure + WHERE drug_concept_id IN (222) + ), + with_prev_max AS ( + SELECT *, + MAX(padded_end) OVER ( + PARTITION BY person_id ORDER BY start_date, padded_end DESC + ROWS BETWEEN UNBOUNDED PRECEDING AND 1 PRECEDING + ) AS prev_max + FROM exposures + ), + with_markers AS ( + SELECT *, + CASE WHEN prev_max IS NULL OR prev_max < start_date THEN 1 ELSE 0 END AS is_new + FROM with_prev_max + ), + with_era AS ( + SELECT *, + SUM(is_new) OVER ( + PARTITION BY person_id + ORDER BY start_date, is_new DESC, padded_end DESC + ) AS era_id + FROM with_markers + ) + SELECT + person_id, + MIN(start_date)::DATE AS era_start_date, + (MAX(padded_end) - {gap})::DATE AS era_end_date + FROM with_era + GROUP BY person_id, era_id + ORDER BY person_id, MIN(start_date) + """ + + raw_conn = conn.con + sql_result = raw_conn.sql(sql).fetchdf() + + # --- compare --- + pd = pytest.importorskip("pandas") + pd.testing.assert_frame_equal( + ibis_result, + sql_result, + check_dtype=False, + check_column_type=False, + ) + + +def test_custom_era_offset_affects_era_grouping(): + """Offset included in padded_end changes which exposures merge into eras. + + With gap_days=0, offset=30: + exp1: end=2020-01-10, exp2: start=2020-01-12 (gap=2 days) + + Without offset in padded_end: padded_end1=2020-01-10 (< start 2020-01-12) + → separate eras, cohort end=2020-01-10+30=2020-02-09 + + With offset in padded_end (Circe BE: DATEADD(day, gap+offset, end)): + padded_end1=2020-01-10+30=2020-02-09 (>= start 2020-01-12) + → merged era, cohort end=2020-01-20+30=2020-02-19 + """ + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1], + "drug_exposure_id": [1, 2], + "drug_concept_id": [222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 1, 12)], + "drug_exposure_end_date": [date(2020, 1, 10), date(2020, 1, 20)], + "days_supply": [0, 0], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 1)], + "condition_end_date": [date(2020, 1, 1)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=0, offset=30), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 1 + assert str(result.iloc[0]["start_date"])[:10] == "2020-01-01" + # Both exposures merge because padded_end1=2020-02-09 >= start 2020-01-12 + # era end = max(end) + offset = 2020-01-20 + 30 = 2020-02-19 + assert str(result.iloc[0]["end_date"])[:10] == "2020-02-19" + + +def test_full_cohort_custom_era_matches_sql_end_dates(): + """Full cohort pipeline with CustomEraStrategy produces same end_dates as raw SQL.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 1], + "drug_exposure_id": [1, 2], + "drug_concept_id": [222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 2, 1)], + "drug_exposure_end_date": [date(2020, 1, 31), date(2020, 3, 3)], + "days_supply": [0, 0], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [100], + "condition_concept_id": [111], + "condition_start_date": [date(2020, 1, 1)], + "condition_end_date": [date(2020, 1, 1)], + "visit_occurrence_id": [10], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + end_strategy=CustomEraStrategy(drug_codeset_id=2, gap_days=30, offset=0), + ) + + # --- ibis pipeline --- + cohort_result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + # --- raw SQL pipeline (Java CUSTOM_ERA_STRATEGY_TEMPLATE logic, DuckDB dialect) --- + # Mirrors Circe BE's generateCohort.sql end-date selection: + # ROW_NUMBER() PARTITION BY person_id, event_id ORDER BY era_end_date ASC + # picks the earliest strategy end per event, matching Circe BE's + # MIN(end_date) across #strategy_ends union. + gap = 30 + sql = f""" + WITH drug_eras AS ( + SELECT + person_id, + MIN(start_date) AS era_start_date, + MAX(padded_end) - {gap} AS era_end_date + FROM ( + SELECT + person_id, start_date, padded_end, + SUM(is_new) OVER ( + PARTITION BY person_id + ORDER BY start_date, is_new DESC, padded_end DESC + ) AS era_id + FROM ( + SELECT + person_id, start_date, padded_end, + CASE WHEN prev_max IS NULL OR prev_max < start_date THEN 1 ELSE 0 END AS is_new + FROM ( + SELECT + person_id, start_date, padded_end, + MAX(padded_end) OVER ( + PARTITION BY person_id ORDER BY start_date, padded_end DESC + ROWS BETWEEN UNBOUNDED PRECEDING AND 1 PRECEDING + ) AS prev_max + FROM ( + SELECT + de.person_id, + de.drug_exposure_start_date::DATE AS start_date, + COALESCE( + de.drug_exposure_end_date::DATE, + de.drug_exposure_start_date::DATE + de.days_supply::INTEGER, + de.drug_exposure_start_date::DATE + 1 + ) + {gap} AS padded_end + FROM drug_exposure de + WHERE de.drug_concept_id = 222 + ) raw_ends + ) maxes + ) marked + ) indexed + GROUP BY person_id, era_id + ), + events_with_obs AS ( + SELECT + e.person_id, + e.condition_occurrence_id AS event_id, + e.condition_start_date::DATE AS start_date, + op.observation_period_end_date::DATE AS op_end_date + FROM condition_occurrence e + JOIN observation_period op ON e.person_id = op.person_id + ), + ranked_ends AS ( + SELECT + ev.person_id, + ev.event_id, + ev.start_date, + ev.op_end_date, + er.era_end_date, + ROW_NUMBER() OVER ( + PARTITION BY ev.person_id, ev.event_id + ORDER BY er.era_end_date + ) AS rn + FROM events_with_obs ev + LEFT JOIN drug_eras er + ON ev.person_id = er.person_id + AND ev.start_date BETWEEN er.era_start_date AND er.era_end_date + ) + SELECT + person_id, + start_date, + LEAST( + COALESCE(era_end_date, op_end_date), + op_end_date + )::DATE AS end_date + FROM ranked_ends + WHERE rn = 1 + ORDER BY person_id, start_date + """ + + sql_result = conn.con.sql(sql).fetchdf() + + # Compare end_dates and start_dates after sorting + ibis_ends = sorted(cohort_result["end_date"].astype(str).tolist()) + sql_ends = sorted(sql_result["end_date"].astype(str).tolist()) + assert ibis_ends == sql_ends + + ibis_starts = sorted(cohort_result["start_date"].astype(str).tolist()) + sql_starts = sorted(sql_result["start_date"].astype(str).tolist()) + assert ibis_starts == sql_starts + + +# --------------------------------------------------------------------------- +# Regression: CustomEra must preserve all events when event_id is shared +# +# After ``first=True`` + ``QualifiedLimit=First`` + ``ExpressionLimit=First`` +# every person contributes at most one event, and ``_assign_primary_event_ids`` +# assigns ``event_id=1`` to all of them. The CustomEra window that selects +# one matching era per event must therefore partition on *(person_id, event_id)* +# — otherwise all rows collapse into a single partition and only one survives. +# --------------------------------------------------------------------------- + + +def _seed_common_tables_multi_person(conn, ibis): + conn.create_table( + "person", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "year_of_birth": [1980, 1985, 1990], + "gender_concept_id": [8507, 8507, 8507], + } + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "observation_period_id": [10, 11, 12], + "observation_period_start_date": [date(2019, 1, 1), date(2019, 1, 1), date(2019, 1, 1)], + "observation_period_end_date": [date(2021, 12, 31), date(2021, 12, 31), date(2021, 12, 31)], + } + ), + overwrite=True, + ) + + +def test_custom_era_preserves_all_persons_with_first_true(): + """All persons survive when DrugExposure(first=True) + CustomEra + limits. + + The window ``group_by=joined.event_id`` previously collapsed every row + into a single partition because all events had ``event_id=1`` (assigned + by ``_assign_primary_event_ids`` — each person has exactly 1 event after + ``first=True`` and the per-person limits). + """ + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables_multi_person(conn, ibis) + + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [1, 2, 3], + "drug_exposure_id": [100, 200, 300], + "drug_concept_id": [222, 222, 222], + "drug_exposure_start_date": [date(2020, 1, 1), date(2020, 2, 1), date(2020, 3, 1)], + "drug_exposure_end_date": [date(2020, 1, 31), date(2020, 2, 28), date(2020, 3, 31)], + "days_supply": [0, 0, 0], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[_make_concept_set(1, 222)], + primary_criteria=PrimaryCriteria(criteria_list=[DrugExposure(codeset_id=1, first=True)]), + qualified_limit=ResultLimit(Type="First"), + expression_limit=ResultLimit(Type="First"), + end_strategy=CustomEraStrategy(drug_codeset_id=1, gap_days=30, offset=0), + ) + + result = build_cohort(expression, backend=conn, cdm_schema="main").execute() + + assert len(result) == 3, f"expected 3 rows, got {len(result)}" + assert set(result["person_id"]) == {1, 2, 3} diff --git a/tests/execution/test_databricks_compat.py b/tests/execution/test_databricks_compat.py index 448bd3df..11d8b220 100644 --- a/tests/execution/test_databricks_compat.py +++ b/tests/execution/test_databricks_compat.py @@ -16,7 +16,8 @@ class FakeDatabricksBackend: def _post_connect(self): raise RuntimeError("CREATE VOLUME IF NOT EXISTS my_catalog.my_schema.memtable") - patched = apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) + with pytest.warns(DeprecationWarning): + patched = apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) assert patched is True backend = FakeDatabricksBackend() @@ -29,7 +30,8 @@ def _post_connect(self): _ = "CREATE VOLUME IF NOT EXISTS my_catalog.my_schema.memtable" raise RuntimeError("different setup error") - patched = apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) + with pytest.warns(DeprecationWarning): + patched = apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) assert patched is True backend = FakeDatabricksBackend() @@ -76,6 +78,8 @@ class FakeDatabricksBackend: def _post_connect(self): raise RuntimeError("CREATE VOLUME IF NOT EXISTS my_catalog.my_schema.memtable") - assert apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) is True + with pytest.warns(DeprecationWarning): + assert apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) is True + # Second call is idempotent — no warning because the patch flag is already set assert apply_databricks_post_connect_workaround(backend_cls=FakeDatabricksBackend) is True assert maybe_apply_databricks_post_connect_workaround(FakeDatabricksBackend()) is True diff --git a/tests/execution/test_error_messages.py b/tests/execution/test_error_messages.py index 80133b45..8e71481b 100644 --- a/tests/execution/test_error_messages.py +++ b/tests/execution/test_error_messages.py @@ -14,7 +14,7 @@ Occurrence, PrimaryCriteria, ) -from circe.cohortdefinition.core import CustomEraStrategy, NumericRange +from circe.cohortdefinition.core import NumericRange from circe.execution.errors import CompilationError, UnsupportedCriterionError, UnsupportedFeatureError from circe.execution.normalize.criteria import normalize_criterion from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem @@ -55,16 +55,6 @@ def _concept_set(set_id: int, concept_id: int) -> ConceptSet: ) -def test_error_message_for_custom_era_end_strategy(): - expression = CohortExpression( - primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence()]), - end_strategy=CustomEraStrategy(drug_codeset_id=1, gap_days=30, offset=0), - ) - - with pytest.raises(UnsupportedFeatureError, match="custom_era end strategy"): - _ = build_cohort(expression, backend=object(), cdm_schema="main") - - def test_error_message_for_unsupported_criterion_type(): with pytest.raises( UnsupportedCriterionError, diff --git a/tests/execution/test_group_demographics.py b/tests/execution/test_group_demographics.py index d11cc730..0fbcb0f4 100644 --- a/tests/execution/test_group_demographics.py +++ b/tests/execution/test_group_demographics.py @@ -1,5 +1,6 @@ from __future__ import annotations +import ibis import pytest from circe.execution.engine.group_demographics import ( @@ -20,8 +21,9 @@ def __init__(self, conn, *, codesets: dict[int, tuple[int, ...]] | None = None): def table(self, name: str): return self.conn.table(name) - def concept_ids_for_codeset(self, codeset_id: int) -> tuple[int, ...]: - return self.codesets.get(codeset_id, ()) + def concept_set_table(self, codeset_id: int) -> ibis.Table: + ids = self.codesets.get(codeset_id, ()) + return ibis.memtable({"concept_id": list(ids)}, schema={"concept_id": "int64"}) def _seed_demographic_tables(conn, ibis): @@ -68,10 +70,22 @@ def test_apply_date_predicate_rejects_invalid_between_and_op(): ) -def test_demographic_concept_ids_merge_codesets_without_duplicates(): - ctx = _DemographicContext(None, codesets={1: (8507, 8532)}) - - assert _demographic_concept_ids(explicit_ids=(8507,), codeset_id=1, ctx=ctx) == (8507, 8532) +def test_demographic_concept_table_returns_table(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + conn = ibis.duckdb.connect() + ctx = _DemographicContext(conn, codesets={1: (8507, 8532)}) + result = _demographic_concept_ids(explicit_ids=(8507,), codeset_id=1, ctx=ctx) + assert result is not None + conn.create_table("_test_demo_concepts", result, temp=True, overwrite=True) + rows = conn.table("_test_demo_concepts").execute() + assert sorted(rows["concept_id"].tolist()) == [8507, 8532] + + +def test_demographic_concept_table_returns_empty_when_empty(): + ctx = _DemographicContext(None) + result = _demographic_concept_ids(explicit_ids=(), codeset_id=None, ctx=ctx) + assert result is None def test_demographic_match_keys_applies_all_supported_filters(): diff --git a/tests/execution/test_inclusion.py b/tests/execution/test_inclusion.py index 46ce20ee..4a63bcff 100644 --- a/tests/execution/test_inclusion.py +++ b/tests/execution/test_inclusion.py @@ -12,6 +12,7 @@ Occurrence, PrimaryCriteria, ) +from circe.execution.api import build_cohort as build_execution_cohort from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem @@ -153,3 +154,112 @@ def test_inclusion_rule_without_expression_is_noop(): result = build_cohort(expression, backend=conn, cdm_schema="main").execute() assert set(result.person_id) == {1, 2} + + +def test_inclusion_rules_materialize_false_matches_materialized_result(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_common_tables(conn, ibis, persons=(1, 2, 3, 4)) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 1, 1, 2, 2, 3, 3, 3, 4, 4], + "condition_occurrence_id": [100, 101, 102, 200, 201, 300, 301, 302, 400, 401], + "condition_concept_id": [111, 222, 333, 111, 222, 111, 222, 333, 111, 333], + "condition_start_date": [ + "2020-01-01", + "2020-01-02", + "2020-01-03", + "2020-01-01", + "2020-01-02", + "2020-01-01", + "2020-01-02", + "2020-01-03", + "2020-01-01", + "2020-01-02", + ], + "condition_end_date": [ + "2020-01-01", + "2020-01-02", + "2020-01-03", + "2020-01-01", + "2020-01-02", + "2020-01-01", + "2020-01-02", + "2020-01-03", + "2020-01-01", + "2020-01-02", + ], + "visit_occurrence_id": [10, 10, 10, 20, 20, 30, 30, 30, 40, 40], + } + ), + overwrite=True, + ) + + expression = CohortExpression( + concept_sets=[ + _make_concept_set(1, 111), + _make_concept_set(2, 222), + _make_concept_set(3, 333), + ], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=1)]), + inclusion_rules=[ + InclusionRule( + name="rule-1", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ConditionOccurrence(codeset_id=2), + occurrence=Occurrence(type=Occurrence._AT_LEAST, count=1), + ) + ], + ), + ), + InclusionRule( + name="rule-2", + expression=CriteriaGroup( + type="ALL", + criteria_list=[ + CorelatedCriteria( + criteria=ConditionOccurrence(codeset_id=3), + occurrence=Occurrence(type=Occurrence._AT_LEAST, count=1), + ) + ], + ), + ), + InclusionRule(name="noop", expression=None), + ], + ) + + materialized = build_execution_cohort( + expression, + backend=conn, + cdm_schema="main", + materialize=True, + ).execute() + non_materialized = build_execution_cohort( + expression, + backend=conn, + cdm_schema="main", + materialize=False, + ).execute() + + materialized_rows = { + tuple(row) + for row in materialized[["person_id", "event_id", "start_date", "end_date"]].itertuples( + index=False, name=None + ) + } + non_materialized_rows = { + tuple(row) + for row in non_materialized[["person_id", "event_id", "start_date", "end_date"]].itertuples( + index=False, name=None + ) + } + + assert non_materialized_rows == materialized_rows + assert materialized_rows diff --git a/tests/execution/test_operations.py b/tests/execution/test_operations.py index ce42448a..2265265c 100644 --- a/tests/execution/test_operations.py +++ b/tests/execution/test_operations.py @@ -96,6 +96,14 @@ def __ne__(self, other): return ("ne", other) +class _CountExpr: + def __init__(self, rows): + self._rows = rows + + def execute(self): + return len(self._rows) + + class _CohortRelation: cohort_definition_id = _CohortColumn() @@ -111,6 +119,9 @@ def filter(self, _predicate): def limit(self, _count): return self + def count(self): + return _CountExpr(self.rows) + def execute(self): return self.rows diff --git a/tests/execution/test_person_filters.py b/tests/execution/test_person_filters.py index 2bb40cdb..0fe018ab 100644 --- a/tests/execution/test_person_filters.py +++ b/tests/execution/test_person_filters.py @@ -1,5 +1,6 @@ from __future__ import annotations +import ibis import pytest from circe.execution.errors import CompilationError @@ -21,8 +22,9 @@ def __init__(self, conn, *, codesets: dict[int, tuple[int, ...]] | None = None): def table(self, name: str): return self.conn.table(name) - def concept_ids_for_codeset(self, codeset_id: int) -> tuple[int, ...]: - return self.codesets.get(codeset_id, ()) + def concept_set_table(self, codeset_id: int): + ids = self.codesets.get(codeset_id, ()) + return ibis.memtable({"concept_id": list(ids)}, schema={"concept_id": "int64"}) def _seed_person_tables(conn, ibis): diff --git a/tests/execution/test_phenotype_failures.py b/tests/execution/test_phenotype_failures.py new file mode 100644 index 00000000..c68e9cc5 --- /dev/null +++ b/tests/execution/test_phenotype_failures.py @@ -0,0 +1,262 @@ +"""Regression tests for PhenotypeLibrary cohorts that previously failed. + +These are the 3 most complex cohorts in the PhenotypeLibrary (51-97 primary +criteria, 105 concept sets, 7-9 inclusion rules, multiple censoring criteria). +They failed because the sequential ``_union_all`` produced deeply nested +UNION ALL expressions that exceeded DuckDB's query compilation limits. + +The binary-tree merge in ``_union_all`` reduces nesting from O(n) to O(log n). + +These tests verify that ``build_cohort`` (compilation) succeeds. Full +``generate_cohort_set`` is covered by ``pytest.mark.slow`` tests. +""" + +from __future__ import annotations + +import datetime +from pathlib import Path + +import pytest + +from circe.cohortdefinition import CohortExpression +from circe.execution.api import build_cohort + +BENCHMARK_OUTPUT = Path(__file__).resolve().parent.parent.parent / "benchmark_output" +JSON_DIR = BENCHMARK_OUTPUT / "phenotype_jsons" + +D = datetime.date # shorthand + + +def _seed_minimal_cdm(conn, ibis): + """Minimal CDM tables so cohort compilation doesn't fail on missing tables. + Uses sentinel person_id=999 so no actual rows match real cohort criteria.""" + S = D(2000, 1, 1) # sentinel date + conn.create_table( + "person", + obj=ibis.memtable({"person_id": [999], "year_of_birth": [1900], "gender_concept_id": [0]}), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [999], + "observation_period_id": [999], + "observation_period_start_date": [S], + "observation_period_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [999], + "condition_occurrence_id": [999], + "condition_concept_id": [0], + "condition_start_date": [S], + "condition_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "procedure_occurrence", + obj=ibis.memtable( + { + "person_id": [999], + "procedure_occurrence_id": [999], + "procedure_concept_id": [0], + "procedure_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "measurement", + obj=ibis.memtable( + { + "person_id": [999], + "measurement_id": [999], + "measurement_concept_id": [0], + "measurement_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "observation", + obj=ibis.memtable( + { + "person_id": [999], + "observation_id": [999], + "observation_concept_id": [0], + "observation_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "drug_exposure", + obj=ibis.memtable( + { + "person_id": [999], + "drug_exposure_id": [999], + "drug_concept_id": [0], + "drug_exposure_start_date": [S], + "drug_exposure_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table("death", obj=ibis.memtable({"person_id": [999], "death_date": [S]}), overwrite=True) + conn.create_table( + "visit_occurrence", + obj=ibis.memtable( + { + "person_id": [999], + "visit_occurrence_id": [999], + "visit_concept_id": [0], + "visit_start_date": [S], + "visit_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "specimen", + obj=ibis.memtable( + {"person_id": [999], "specimen_id": [999], "specimen_concept_id": [0], "specimen_date": [S]} + ), + overwrite=True, + ) + conn.create_table( + "device_exposure", + obj=ibis.memtable( + { + "person_id": [999], + "device_exposure_id": [999], + "device_concept_id": [0], + "device_exposure_start_date": [S], + "device_exposure_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "dose_era", + obj=ibis.memtable( + { + "person_id": [999], + "dose_era_id": [999], + "drug_concept_id": [0], + "unit_concept_id": [0], + "dose_value": [0.0], + "dose_era_start_date": [S], + "dose_era_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "payer_plan_period", + obj=ibis.memtable( + { + "person_id": [999], + "payer_plan_period_id": [999], + "payer_plan_period_start_date": [S], + "payer_plan_period_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "visit_detail", + obj=ibis.memtable( + { + "person_id": [999], + "visit_detail_id": [999], + "visit_detail_concept_id": [0], + "visit_detail_start_date": [S], + "visit_detail_end_date": [S], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_era", + obj=ibis.memtable( + { + "person_id": [999], + "condition_era_id": [999], + "condition_concept_id": [0], + "condition_era_start_date": [S], + "condition_era_end_date": [S], + "condition_occurrence_count": [1], + } + ), + overwrite=True, + ) + conn.create_table( + "drug_era", + obj=ibis.memtable( + { + "person_id": [999], + "drug_era_id": [999], + "drug_concept_id": [0], + "drug_era_start_date": [S], + "drug_era_end_date": [S], + "drug_exposure_count": [1], + "gap_days": [0], + } + ), + overwrite=True, + ) + conn.create_table( + "concept", obj=ibis.memtable({"concept_id": [0, 999], "invalid_reason": ["X", None]}), overwrite=True + ) + conn.create_table( + "concept_ancestor", + obj=ibis.memtable({"ancestor_concept_id": [999], "descendant_concept_id": [999]}), + overwrite=True, + ) + conn.create_table( + "concept_relationship", + obj=ibis.memtable( + {"concept_id_1": [999], "concept_id_2": [999], "relationship_id": ["X"], "invalid_reason": ["X"]} + ), + overwrite=True, + ) + + +# The 3 persistently failing PhenotypeLibrary cohorts +FAILING_COHORT_IDS = [1432, 1433, 1434] + + +@pytest.mark.parametrize("cohort_id", FAILING_COHORT_IDS) +def test_phenotype_cohort_compiles(cohort_id: int) -> None: + """Cohorts 1432-1434 (87, 97, 51 primary criteria) must compile successfully. + + Compilation exercises ``_union_all`` which previously produced O(n) nested + UNION ALL expressions that crashed DuckDB. The binary-tree merge reduces + nesting to O(log n). + """ + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + json_path = JSON_DIR / f"{cohort_id}.json" + if not json_path.exists(): + pytest.skip(f"Phenotype JSON not found: {json_path}") + + expression = CohortExpression.model_validate_json(json_path.read_text()) + + conn = ibis.duckdb.connect() + _seed_minimal_cdm(conn, ibis) + + # build_cohort exercises the full UNION ALL path; materialize=False + # keeps this compile-only (no temp tables created). + try: + build_cohort(expression, backend=conn, cdm_schema="main", materialize=False) + except Exception as exc: + pytest.fail(f"Cohort {cohort_id} compilation failed: {exc}") diff --git a/tests/execution/test_registry_dispatch.py b/tests/execution/test_registry_dispatch.py index 7d53701a..ca1c98f5 100644 --- a/tests/execution/test_registry_dispatch.py +++ b/tests/execution/test_registry_dispatch.py @@ -40,6 +40,7 @@ def test_registry_dispatch_round_trip(): source_concept_column=None, visit_occurrence_column=None, codeset_id=None, + source_codeset_id=None, first=False, occurrence_start_date=None, occurrence_end_date=None, @@ -87,6 +88,7 @@ class UnknownCriteria(Criteria): source_concept_column=None, visit_occurrence_column=None, codeset_id=None, + source_codeset_id=None, first=False, occurrence_start_date=None, occurrence_end_date=None, diff --git a/tests/execution/test_union_scaling.py b/tests/execution/test_union_scaling.py new file mode 100644 index 00000000..175ec311 --- /dev/null +++ b/tests/execution/test_union_scaling.py @@ -0,0 +1,119 @@ +"""Tests that cohort expression UNION ALL scales to large numbers of primary criteria. + +Reproduces the failure mode where cohorts with 51-97 criteria (like +PhenotypeLibrary cohorts 1432-1434) crash DuckDB because the sequential +pairwise ``_union_all`` produces O(n) nesting depth. +""" + +from __future__ import annotations + +import pytest + +from circe.cohort_definition_set import CohortDefinitionSet, generate_cohort_set +from circe.cohortdefinition import CohortExpression, ConditionOccurrence, PrimaryCriteria +from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem + + +def _make_concept_set(set_id: int, concept_id: int) -> ConceptSet: + return ConceptSet( + id=set_id, + expression=ConceptSetExpression(items=[ConceptSetItem(concept=Concept(conceptId=concept_id))]), + ) + + +def _seed_tables(conn, ibis, n_persons: int = 2): + person_ids = list(range(1, n_persons + 1)) + conn.create_table( + "person", + obj=ibis.memtable( + { + "person_id": person_ids, + "year_of_birth": [1980] * n_persons, + "gender_concept_id": [8507] * n_persons, + } + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": person_ids, + "observation_period_id": person_ids, + "observation_period_start_date": ["2019-01-01"] * n_persons, + "observation_period_end_date": ["2022-12-31"] * n_persons, + } + ), + overwrite=True, + ) + + # Each person gets one condition row per concept_id from 1..N + condition_rows = [] + for pid in person_ids: + for cid in range(1, n_persons * 25 + 1): + condition_rows.append( + { + "person_id": pid, + "condition_occurrence_id": pid * 1000 + cid, + "condition_concept_id": cid, + "condition_start_date": "2020-01-10", + "condition_end_date": "2020-01-10", + } + ) + + conn.create_table("condition_occurrence", obj=ibis.memtable(condition_rows), overwrite=True) + + # Minimal vocabulary (must have at least one non-None invalid_reason + # so ibis can infer a non-NULL column type for DuckDB) + concept_ids = list(range(1, n_persons * 25 + 1)) + invalid = [None] * len(concept_ids) + if invalid: + invalid[0] = "X" # ensure type inference + conn.create_table( + "concept", + obj=ibis.memtable( + { + "concept_id": concept_ids, + "invalid_reason": invalid, + } + ), + overwrite=True, + ) + + +def _build_multi_criterion_expression(n: int) -> CohortExpression: + """Build a cohort with *n* simple ConditionOccurrence primary criteria.""" + concept_sets = [] + criteria = [] + for i in range(n): + cs_id = i + 1 + concept_sets.append(_make_concept_set(cs_id, concept_id=cs_id)) + criteria.append(ConditionOccurrence(codeset_id=cs_id)) + return CohortExpression( + concept_sets=concept_sets, primary_criteria=PrimaryCriteria(criteria_list=criteria) + ) + + +@pytest.mark.parametrize("n_criteria", [1, 2, 5, 10, 20, 50, 100]) +def test_union_all_scales(n_criteria: int) -> None: + """Cohorts with N primary criteria should compile and execute without nesting errors.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + expression = _build_multi_criterion_expression(n_criteria) + cds = CohortDefinitionSet() + cds.add(1, f"UnionTest_{n_criteria}", expression) + + results = generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table=f"union_test_{n_criteria}", + stop_on_error=True, + ) + + assert len(results) == 1 + assert results[0].status == "COMPLETE", f"{n_criteria} criteria failed: {results[0].error}" diff --git a/tests/test_cohort_definition_set.py b/tests/test_cohort_definition_set.py new file mode 100644 index 00000000..9b7e3438 --- /dev/null +++ b/tests/test_cohort_definition_set.py @@ -0,0 +1,775 @@ +"""Tests for CohortDefinitionSet and generate_cohort_set.""" + +from __future__ import annotations + +import pytest + +from circe.cohort_definition_set import ( + CohortDefinition, + CohortDefinitionSet, + CohortGenerationResult, + generate_cohort_set, + summarise_generation_results, +) +from circe.cohortdefinition import CohortExpression, ConditionOccurrence, PrimaryCriteria +from circe.vocabulary import Concept, ConceptSet, ConceptSetExpression, ConceptSetItem + + +def _make_concept_set(set_id: int, concept_id: int) -> ConceptSet: + return ConceptSet( + id=set_id, + expression=ConceptSetExpression(items=[ConceptSetItem(concept=Concept(conceptId=concept_id))]), + ) + + +def _simple_expression(concept_id: int = 111, set_id: int = 1) -> CohortExpression: + return CohortExpression( + concept_sets=[_make_concept_set(set_id, concept_id)], + primary_criteria=PrimaryCriteria(criteria_list=[ConditionOccurrence(codeset_id=set_id)]), + ) + + +def _seed_tables(conn, ibis): + conn.create_table( + "person", + obj=ibis.memtable( + { + "person_id": [1, 2], + "year_of_birth": [1980, 1982], + "gender_concept_id": [8507, 8507], + } + ), + overwrite=True, + ) + conn.create_table( + "observation_period", + obj=ibis.memtable( + { + "person_id": [1, 2], + "observation_period_id": [10, 11], + "observation_period_start_date": ["2019-01-01", "2019-01-01"], + "observation_period_end_date": ["2021-12-31", "2021-12-31"], + } + ), + overwrite=True, + ) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1, 2], + "condition_occurrence_id": [100, 101], + "condition_concept_id": [111, 111], + "condition_start_date": ["2020-01-01", "2020-01-02"], + "condition_end_date": ["2020-01-01", "2020-01-02"], + } + ), + overwrite=True, + ) + + +# --------------------------------------------------------------------------- +# CohortDefinitionSet — unit tests (no database needed) +# --------------------------------------------------------------------------- + + +def test_cohort_definition_set_add_and_iter(): + expr = _simple_expression() + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Cohort A", expression=expr) + cds.add(cohort_id=2, cohort_name="Cohort B", expression=expr) + + assert len(cds) == 2 + ids = [c.cohort_id for c in cds] + assert ids == [1, 2] + + +def test_cohort_definition_set_duplicate_id_raises(): + expr = _simple_expression() + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Cohort A", expression=expr) + with pytest.raises(ValueError, match="cohort_id=1"): + cds.add(cohort_id=1, cohort_name="Duplicate", expression=expr) + + +def test_cohort_definition_set_getitem(): + expr = _simple_expression() + cds = CohortDefinitionSet() + cds.add(cohort_id=42, cohort_name="My Cohort", expression=expr) + + item = cds[42] + assert isinstance(item, CohortDefinition) + assert item.cohort_id == 42 + assert item.cohort_name == "My Cohort" + + +def test_cohort_definition_set_getitem_missing_raises(): + cds = CohortDefinitionSet() + with pytest.raises(KeyError): + _ = cds[999] + + +def test_checksums_returns_dict(): + expr1 = _simple_expression(concept_id=111) + expr2 = _simple_expression(concept_id=222, set_id=2) + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=expr1) + cds.add(cohort_id=2, cohort_name="B", expression=expr2) + + checksums = cds.checksums() + assert set(checksums.keys()) == {1, 2} + assert isinstance(checksums[1], str) + assert len(checksums[1]) == 64 # SHA-256 hex digest + # Different expressions produce different checksums + assert checksums[1] != checksums[2] + + +def test_checksums_stable(): + expr = _simple_expression() + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=expr) + + assert cds.checksums()[1] == cds.checksums()[1] + + +# --------------------------------------------------------------------------- +# generate_cohort_set — integration tests using DuckDB +# --------------------------------------------------------------------------- + + +def test_generate_cohort_set_basic(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=10, cohort_name="Cohort 10", expression=_simple_expression()) + cds.add(cohort_id=20, cohort_name="Cohort 20", expression=_simple_expression()) + + results = generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="cohort_out") + + assert len(results) == 2 + assert all(r.status == "COMPLETE" for r in results) + assert {r.cohort_id for r in results} == {10, 20} + + cohort_table = conn.table("cohort_out").execute() + assert len(cohort_table) == 4 # 2 persons × 2 cohorts + assert set(cohort_table.cohort_definition_id) == {10, 20} + + +def test_generate_cohort_set_results_have_timing(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + + results = generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="c") + + r = results[0] + assert r.start_time <= r.end_time + assert isinstance(r.checksum, str) and len(r.checksum) == 64 + + +def test_generate_cohort_set_incremental_skip(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="B", expression=_simple_expression()) + + # First run: both should be COMPLETE + first = generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=True) + assert all(r.status == "COMPLETE" for r in first) + + # Second run with same expressions: both should be SKIPPED + second = generate_cohort_set( + cds, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=True + ) + assert all(r.status == "SKIPPED" for r in second) + + +def test_generate_cohort_set_incremental_regenerate(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + expr_a = _simple_expression(concept_id=111) + expr_b = _simple_expression(concept_id=111) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=expr_a) + cds.add(cohort_id=2, cohort_name="B", expression=expr_b) + + generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=True) + + # Change cohort 1's expression (different concept) + conn.create_table( + "condition_occurrence", + obj=ibis.memtable( + { + "person_id": [1], + "condition_occurrence_id": [200], + "condition_concept_id": [222], + "condition_start_date": ["2020-03-01"], + "condition_end_date": ["2020-03-01"], + } + ), + overwrite=True, + ) + expr_a_changed = _simple_expression(concept_id=222) # different concept + + cds2 = CohortDefinitionSet() + cds2.add(cohort_id=1, cohort_name="A", expression=expr_a_changed) + cds2.add(cohort_id=2, cohort_name="B", expression=expr_b) # unchanged + + results = generate_cohort_set( + cds2, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=True + ) + + statuses = {r.cohort_id: r.status for r in results} + assert statuses[1] == "COMPLETE" # regenerated + assert statuses[2] == "SKIPPED" # unchanged + + +def test_generate_cohort_set_incremental_non_incremental_does_not_skip(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + + generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=True) + + # Non-incremental run should always be COMPLETE regardless of stored checksums + results = generate_cohort_set( + cds, backend=conn, cdm_schema="main", cohort_table="cohort", incremental=False + ) + assert results[0].status == "COMPLETE" + + +def test_generate_cohort_set_continue_on_error(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + # Use a CohortExpression with a concept set referencing a missing table domain + # by using a bad backend-level call; we monkeypatch write_cohort instead. + from unittest.mock import patch + + from circe.execution.errors import ExecutionError + + call_count = 0 + + def _failing_write_cohort(*, compiled_relation, cohort_id, **kwargs): + nonlocal call_count + call_count += 1 + if cohort_id == 1: + raise ExecutionError("Simulated failure for cohort 1") + # Delegate to real write_cohort for cohort 2 + from circe.execution.api import write_cohort as real_write_cohort + + real_write_cohort(compiled_relation=compiled_relation, cohort_id=cohort_id, **kwargs) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Good", expression=_simple_expression()) + + with patch("circe.cohort_definition_set._generate.write_cohort", side_effect=_failing_write_cohort): + results = generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort", + stop_on_error=False, + ) + + assert call_count == 2 + statuses = {r.cohort_id: r.status for r in results} + assert statuses[1] == "FAILED" + assert statuses[2] == "COMPLETE" + assert results[0].error is not None + + +def test_generate_cohort_set_continue_on_non_execution_error(): + """Non-ExecutionError exceptions must also be caught and recorded.""" + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from unittest.mock import patch + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + call_count = 0 + + def _failing_build(expression, *, backend, cohort_id, **kwargs): + nonlocal call_count + call_count += 1 + if cohort_id == 1: + raise RuntimeError("Simulated RuntimeError") + from circe.execution.api import build_cohort as real_build + + return real_build(expression, backend=backend, cohort_id=cohort_id, **kwargs) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Good", expression=_simple_expression()) + + with patch( + "circe.cohort_definition_set._generate.build_cohort", + side_effect=_failing_build, + ): + results = generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_non_exec", + stop_on_error=False, + ) + + assert call_count == 2 + statuses = {r.cohort_id: r.status for r in results} + assert statuses[1] == "FAILED" + assert statuses[2] == "COMPLETE" + assert isinstance(results[0].error, RuntimeError) + + +def test_generate_cohort_set_stop_on_error(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from unittest.mock import patch + + from circe.execution.errors import ExecutionError + + def _always_fail(*, compiled_relation, cohort_id, **kwargs): + raise ExecutionError("Always fail") + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Also bad", expression=_simple_expression()) + + with ( + patch("circe.cohort_definition_set._generate.write_cohort", side_effect=_always_fail), + pytest.raises(ExecutionError, match="Always fail"), + ): + generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort", + stop_on_error=True, + ) + + +def test_summarise_generation_results(): + from datetime import datetime + + now = datetime.now() + + results = [ + CohortGenerationResult(1, "A", "COMPLETE", "abc", now, now), + CohortGenerationResult(2, "B", "SKIPPED", "def", now, now), + CohortGenerationResult(3, "C", "FAILED", "ghi", now, now), + CohortGenerationResult(4, "D", "COMPLETE", "jkl", now, now), + ] + + summary = summarise_generation_results(results) + assert summary["COMPLETE"] == 2 + assert summary["SKIPPED"] == 1 + assert summary["FAILED"] == 1 + + +# --------------------------------------------------------------------------- +# Generation history table integration tests +# --------------------------------------------------------------------------- + + +def test_generate_cohort_set_history_table_populated(): + import pandas as pd + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="B", expression=_simple_expression()) + + CHECKSUM_TABLE = "cohort_checksum_test" + + results = generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort", + incremental=True, + checksum_table=CHECKSUM_TABLE, + ) + assert all(r.status == "COMPLETE" for r in results) + + history = conn.table(CHECKSUM_TABLE, database="main").execute() + assert not history.empty + assert "cohort_definition_id" in history.columns + assert "checksum" in history.columns + assert "status" in history.columns + assert "start_time" in history.columns + assert "end_time" in history.columns + + for _, row in history.iterrows(): + assert row["status"] in ("COMPLETE", "FAILED") + assert pd.to_datetime(row["end_time"]) >= pd.to_datetime(row["start_time"]) + + assert set(history["cohort_definition_id"]) == {1, 2} + assert all(history["status"] == "COMPLETE") + + +def test_generate_cohort_set_history_table_skip_no_duplicate(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="B", expression=_simple_expression()) + + CHECKSUM_TABLE = "cohort_checksum_skip_test" + + # Run 1: both COMPLETE → both get history entries + generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_skip", + incremental=True, + checksum_table=CHECKSUM_TABLE, + ) + after_first = conn.table(CHECKSUM_TABLE, database="main").execute() + assert len(after_first) == 2 # 2 history rows + + # Run 2: incremental, all should be SKIPPED → no new history entries + generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_skip", + incremental=True, + checksum_table=CHECKSUM_TABLE, + ) + after_second = conn.table(CHECKSUM_TABLE, database="main").execute() + assert len(after_second) == 2 # still 2 — no duplicates for SKIPPED + + +def test_generate_cohort_set_history_table_failed(): + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from unittest.mock import patch + + from circe.execution.errors import ExecutionError + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Good", expression=_simple_expression()) + + CHECKSUM_TABLE = "cohort_checksum_fail_test" + + call_count = 0 + + def _failing_write(*, compiled_relation, cohort_id, **kwargs): + nonlocal call_count + call_count += 1 + if cohort_id == 1: + raise ExecutionError("Simulated failure for cohort 1") + from circe.execution.api import write_cohort as real_write_cohort + + real_write_cohort(compiled_relation=compiled_relation, cohort_id=cohort_id, **kwargs) + + with patch("circe.cohort_definition_set._generate.write_cohort", side_effect=_failing_write): + results = generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_fail", + incremental=True, + checksum_table=CHECKSUM_TABLE, + stop_on_error=False, + ) + + statuses = {r.cohort_id: r.status for r in results} + assert statuses[1] == "FAILED" + assert statuses[2] == "COMPLETE" + + history = conn.table(CHECKSUM_TABLE, database="main").execute() + history_statuses = dict(zip(history["cohort_definition_id"], history["status"], strict=True)) + assert history_statuses[1] == "FAILED" + assert history_statuses[2] == "COMPLETE" + + +def test_load_generation_history(): + from circe.cohort_definition_set._checksum_store import load_generation_history + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + + CHECKSUM_TABLE = "cohort_history_test" + + generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_hist", + incremental=True, + checksum_table=CHECKSUM_TABLE, + ) + + history = load_generation_history(conn, schema="main", table_name=CHECKSUM_TABLE) + assert history is not None + rows = history.execute() + assert not rows.empty + assert "start_time" in rows.columns + assert "end_time" in rows.columns + assert "status" in rows.columns + assert rows.iloc[0]["status"] == "COMPLETE" + + # Non-existent table returns None + none_result = load_generation_history(conn, schema="main", table_name="nonexistent_table") + assert none_result is None + + +def test_api_exports_cohort_definition_set(): + import circe.api as api + + assert hasattr(api, "CohortDefinitionSet") + assert hasattr(api, "CohortDefinition") + assert hasattr(api, "CohortGenerationResult") + assert hasattr(api, "async_generate_cohort_set") + assert hasattr(api, "generate_cohort_set") + assert hasattr(api, "summarise_generation_results") + + +# --------------------------------------------------------------------------- +# async_generate_cohort_set tests +# --------------------------------------------------------------------------- + + +def test_async_generate_cohort_set_basic(): + import asyncio + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from circe.cohort_definition_set._generate import async_generate_cohort_set + + cds = CohortDefinitionSet() + cds.add(cohort_id=10, cohort_name="Cohort 10", expression=_simple_expression()) + cds.add(cohort_id=20, cohort_name="Cohort 20", expression=_simple_expression()) + + results = asyncio.run( + async_generate_cohort_set(cds, backend=conn, cdm_schema="main", cohort_table="cohort_async") + ) + + assert len(results) == 2 + assert all(r.status == "COMPLETE" for r in results) + assert {r.cohort_id for r in results} == {10, 20} + + cohort_table = conn.table("cohort_async").execute() + assert set(cohort_table.cohort_definition_id) == {10, 20} + + +def test_async_generate_cohort_set_continue_on_non_execution_error(): + """Non-ExecutionError exceptions (e.g. databricks ServerOperationError) + must be caught and recorded as FAILED.""" + import asyncio + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from unittest.mock import patch + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from circe.cohort_definition_set._generate import async_generate_cohort_set + + call_count = 0 + + def _failing_build(expression, *, backend, cohort_id, **kwargs): + nonlocal call_count + call_count += 1 + if cohort_id == 1: + raise RuntimeError("Simulated non-ExecutionError failure") + from circe.execution.api import build_cohort as real_build + + return real_build(expression, backend=backend, cohort_id=cohort_id, **kwargs) + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Good", expression=_simple_expression()) + + with patch( + "circe.cohort_definition_set._generate.build_cohort", + side_effect=_failing_build, + ): + results = asyncio.run( + async_generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_non_exec", + stop_on_error=False, + ) + ) + + assert call_count == 2 + statuses = {r.cohort_id: r.status for r in results} + assert statuses[1] == "FAILED" + assert statuses[2] == "COMPLETE" + assert isinstance(results[0].error, RuntimeError) + + +def test_async_generate_cohort_set_stop_on_error(): + import asyncio + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from unittest.mock import patch + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from circe.cohort_definition_set._generate import async_generate_cohort_set + + def _always_fail(expression, *, backend, cohort_id, **kwargs): + raise ValueError("Simulated non-ExecutionError failure") + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Bad", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="Also bad", expression=_simple_expression()) + + with ( + patch( + "circe.cohort_definition_set._generate.build_cohort", + side_effect=_always_fail, + ), + pytest.raises(ValueError, match="Simulated non-ExecutionError failure"), + ): + asyncio.run( + async_generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_stop", + stop_on_error=True, + ) + ) + + +def test_async_generate_cohort_set_timeout(): + import asyncio + import time + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + from unittest.mock import patch + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from circe.cohort_definition_set._generate import async_generate_cohort_set + + def _slow_build(expression, *, backend, cohort_id, **kwargs): + time.sleep(0.5) + raise RuntimeError("should have timed out") + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="Slow", expression=_simple_expression()) + + with patch( + "circe.cohort_definition_set._generate.build_cohort", + side_effect=_slow_build, + ): + results = asyncio.run( + async_generate_cohort_set( + cds, + backend=conn, + cdm_schema="main", + cohort_table="cohort_timeout", + stop_on_error=False, + compile_timeout=0.1, + ) + ) + + assert len(results) == 1 + assert results[0].status == "FAILED" + assert "timeout" in str(results[0].error).lower() + + +def test_async_generate_cohort_set_incremental_skip(): + import asyncio + + ibis = pytest.importorskip("ibis") + _ = pytest.importorskip("duckdb") + + conn = ibis.duckdb.connect() + _seed_tables(conn, ibis) + + from circe.cohort_definition_set._generate import async_generate_cohort_set + + cds = CohortDefinitionSet() + cds.add(cohort_id=1, cohort_name="A", expression=_simple_expression()) + cds.add(cohort_id=2, cohort_name="B", expression=_simple_expression()) + + # First run -- both COMPLETE + first = asyncio.run( + async_generate_cohort_set( + cds, backend=conn, cdm_schema="main", cohort_table="cohort_inc_async", incremental=True + ) + ) + assert all(r.status == "COMPLETE" for r in first) + + # Second run -- both SKIPPED + second = asyncio.run( + async_generate_cohort_set( + cds, backend=conn, cdm_schema="main", cohort_table="cohort_inc_async", incremental=True + ) + ) + assert all(r.status == "SKIPPED" for r in second) diff --git a/tests/test_extension_system.py b/tests/test_extension_system.py index e76c16b0..bdfedf1f 100644 --- a/tests/test_extension_system.py +++ b/tests/test_extension_system.py @@ -1,5 +1,4 @@ import json -from typing import Optional from pydantic import AliasChoices, Field @@ -26,12 +25,12 @@ class WeatherCondition(Criteria): Imagine a CDM extension where weather data is linked to persons. """ - weather_concept_id: Optional[list[Concept]] = Field( + weather_concept_id: list[Concept] | None = Field( default=None, validation_alias=AliasChoices("WeatherConceptId", "weatherConceptId"), serialization_alias="WeatherConceptId", ) - temperature_celsius: Optional[float] = Field( + temperature_celsius: float | None = Field( default=None, validation_alias=AliasChoices("TemperatureCelsius", "temperatureCelsius"), serialization_alias="TemperatureCelsius", diff --git a/tests/test_real_example_cohorts.py b/tests/test_real_example_cohorts.py index cb3fae6b..aaf477ad 100644 --- a/tests/test_real_example_cohorts.py +++ b/tests/test_real_example_cohorts.py @@ -14,7 +14,6 @@ import textwrap from difflib import unified_diff from pathlib import Path -from typing import Optional import pytest @@ -61,7 +60,7 @@ def pytest_generate_tests(metafunc): metafunc.parametrize("cohort_name", params) -def get_reference_sql(cohort_name: str) -> Optional[str]: +def get_reference_sql(cohort_name: str) -> str | None: """Get pre-generated reference SQL from R/Java implementation.""" ref_file = REFERENCE_DIR / cohort_name.replace(".json", ".sql") if ref_file.exists(): @@ -69,7 +68,7 @@ def get_reference_sql(cohort_name: str) -> Optional[str]: return None -def generate_python_outputs(cohort_file: Path) -> tuple[Optional[str], Optional[str]]: +def generate_python_outputs(cohort_file: Path) -> tuple[str | None, str | None]: """ Run Python reference implementation to generate SQL. @@ -392,10 +391,10 @@ def test_sql_matches_reference(cohort_name): # ============================================================================= # Cache for generated markdown to avoid redundant work -_MARKDOWN_CACHE: dict[str, tuple[Optional[str], Optional[str]]] = {} +_MARKDOWN_CACHE: dict[str, tuple[str | None, str | None]] = {} -def get_generated_markdown(cohort_name: str) -> tuple[Optional[str], Optional[str]]: +def get_generated_markdown(cohort_name: str) -> tuple[str | None, str | None]: """ Get generated markdown for a cohort, using cache if available. """ @@ -424,7 +423,7 @@ def get_generated_markdown(cohort_name: str) -> tuple[Optional[str], Optional[st return markdown, error -def get_reference_markdown(cohort_name: str) -> Optional[str]: +def get_reference_markdown(cohort_name: str) -> str | None: """Get pre-generated reference Markdown from R/Java implementation.""" ref_file = REFERENCE_DIR / cohort_name.replace(".json", ".md") if ref_file.exists(): diff --git a/tox.ini b/tox.ini index 2d25e303..75d0cbf7 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = py39, py310, py311, py312, py313, py314 +envlist = py310, py311, py312, py313, py314 skip_missing_interpreters = true isolated_build = true @@ -21,7 +21,6 @@ commands = [gh-actions] python = - 3.9: py39 3.10: py310 3.11: py311 3.12: py312