Merged in feature/DAIP2-1615-refactor-maximize-code-reusability (pull request #930)

Feature/DAIP2-1615 refactor maximize code reusability

* added reusability

* client resolver added

* Merged DEV into feature/DAIP2-1615-refactor-maximize-code-reusability

* Refactor client prompt reuse and harden resolver test coverage.

Remove duplicated client prompt functions so unresolved names fall back to SaaS, add resolver contract tests with runtime-context isolation, and clean resolver-path test imports to prevent cross-test client leakage.

* Merge origin/dev into feature/DAIP2-1615-refactor-maximize-code-reusability.

Bring in latest vendor pipeline framework changes from dev and resolve runner routing conflict by preserving vendor processor dispatch while keeping shared function-level client fallback behavior.

* Merged dev into feature/DAIP2-1615-refactor-maximize-code-reusability

* Remove tracked drift review report from documentation.

Keep drift/alignment notes as local-only artifacts under ignored docs/ and stop tracking documentation/function_level_override_drift_review.md in the branch.

* Merged dev into feature/DAIP2-1615-refactor-maximize-code-reusability


Approved-by: Katon Minhas
This commit is contained in:
Praneel Panchigar
2026-04-06 21:47:27 +00:00
committed by Katon Minhas
parent 3d3c1913ef
commit 5743a4c1be
22 changed files with 378 additions and 353 deletions
+4 -1
View File
@@ -171,7 +171,10 @@ class BasePipeline(ABC):
self, file_path: str, **kwargs
) -> Tuple[pd.DataFrame, Dict]:
"""Default file processing using saas module."""
from src.pipelines.saas import file_processing
from src.pipelines.runtime_context import set_active_client
from src.pipelines.shared import file_processing
set_active_client(self.client_name)
return file_processing.process_file(file_path, **kwargs)
@@ -44,107 +44,6 @@ def prompt_exhibit_level(
return llm_answer_final
def prompt_exhibit_level_breakout(
exhibit_level_answers: dict[str, str], filename: str
) -> dict[str, str]:
"""
Process exhibit-level answers by breaking out facility adjustment terms using an LLM.
This function checks if a FACILITY_ADJUSTMENT_TERM exists in the exhibit-level answers,
and if present, invokes Claude LLM to break it down into structured components. The
resulting parsed data is then merged back into the exhibit-level answers dictionary.
Args:
exhibit_level_answers (dict[str, str]): A dictionary containing exhibit-level
field names as keys and their corresponding values as strings. Must contain
a "FACILITY_ADJUSTMENT_TERM" key to trigger processing.
filename (str): The name of the file being processed, used for logging and
tracking purposes in the LLM invocation.
Returns:
dict[str, str]: The updated exhibit_level_answers dictionary with additional
fields extracted from the FACILITY_ADJUSTMENT_TERM breakout, or the original
dictionary if no FACILITY_ADJUSTMENT_TERM was present.
Raises:
Any exceptions raised by llm_utils.invoke_claude() or _parser()
will propagate to the caller.
"""
# FACILITY_ADJUSTMENT_BREAKOUT
if not string_utils.is_empty(
exhibit_level_answers.get("FACILITY_ADJUSTMENT_TERM", "")
):
prompt, _parser = prompt_templates.FACILITY_ADJUSTMENT_BREAKOUT(
exhibit_level_answers["FACILITY_ADJUSTMENT_TERM"]
)
llm_answer_raw = llm_utils.invoke_claude(
prompt=prompt, model_id="sonnet_latest", filename=filename
)
llm_answer_final = _parser(llm_answer_raw)
exhibit_level_answers.update(llm_answer_final)
return exhibit_level_answers
def prompt_dynamic_primary(
exhibit_text: str, field: Field, constants: Constants, filename: str, TEMPLATE
):
# Check if we should use context caching for DYNAMIC_PRIMARY
if TEMPLATE == prompt_templates.DYNAMIC_PRIMARY:
# Use the context caching version that splits context from field question
context_text, prompt, _parser = prompt_templates.DYNAMIC_PRIMARY(
exhibit_text,
field.field_name,
field.get_prompt(constants),
)
logging.debug(
f"Dynamic primary prompt with context caching for {filename}; {field}: {prompt}"
)
logging.debug(f"Context length for caching: {len(context_text)} chars")
# Invoke with context_for_caching to enable exhibit-level caching
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.DYNAMIC_PRIMARY_INSTRUCTION(),
context_for_caching=context_text,
usage_label="DYNAMIC_PRIMARY",
)
else:
# Support both legacy (prompt, parser) and context-aware
# (context_text, prompt, parser) template return signatures.
template_result = TEMPLATE(
exhibit_text, field.field_name, field.get_prompt(constants)
)
if not isinstance(template_result, (tuple, list)):
raise ValueError(
f"Template must return tuple or list, got {type(template_result)} "
f"for {getattr(TEMPLATE, '__name__', TEMPLATE)}"
)
if len(template_result) == 3:
context_text, prompt, _parser = template_result
elif len(template_result) == 2:
prompt, _parser = template_result
context_text = None
else:
raise ValueError(
f"Unexpected template return size for {getattr(TEMPLATE, '__name__', TEMPLATE)}: {len(template_result)}"
)
logging.debug(f"Dynamic primary prompt for {filename}; {field}: {prompt}")
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.DYNAMIC_PRIMARY_INSTRUCTION(),
context_for_caching=context_text,
usage_label="DYNAMIC_PRIMARY",
)
logging.debug(f"Claude answer for {filename}; {field}: {llm_answer_raw}")
llm_answer_final = _parser(llm_answer_raw)
return llm_answer_final
def prompt_reimbursement_primary(
page_text: str,
filename: str,
@@ -295,37 +194,6 @@ def prompt_fee_schedule_breakout(
return fs_breakout_dict
def prompt_grouper_breakout(
service: str,
reimbursement_method: str,
filename: str,
):
"""
Call GROUPER_BREAKOUT prompt with cached instruction.
Note: Field definitions are now included in GROUPER_BREAKOUT_INSTRUCTION() for caching.
"""
prompt, _parser = prompt_templates.GROUPER_BREAKOUT(
service,
reimbursement_method,
)
logging.debug(f"Grouper Breakout Prompt for {filename}: {prompt}")
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.GROUPER_BREAKOUT_INSTRUCTION(),
usage_label="GROUPER_BREAKOUT",
)
logging.debug(f"LLM Response for {filename}: {llm_answer_raw}")
try:
grouper_breakout_dict = _parser(llm_answer_raw)
except:
grouper_breakout_dict = {}
return grouper_breakout_dict
def prompt_carveout_check(
service_term: str,
reimb_term: str,
@@ -352,39 +220,6 @@ def prompt_carveout_check(
return carveout_answer
def prompt_special_case_breakout(breakout_template, term_answer, filename):
# Get the prompt from the template
prompt, _parser = breakout_template(term_answer)
# Ensure prompt is a string for Bedrock messages API
if not isinstance(prompt, str):
logging.warning(
f"Expected prompt string, got {type(prompt)} in prompt_special_case_breakout; coercing to str."
)
prompt = str(prompt)
# Look for a corresponding _INSTRUCTION() function for caching
template_name = getattr(breakout_template, "__name__", "SPECIAL_CASE_BREAKOUT")
instruction_func_name = f"{template_name}_INSTRUCTION"
instruction_func = getattr(prompt_templates, instruction_func_name, None)
if instruction_func and callable(instruction_func):
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=instruction_func(),
usage_label=template_name,
)
else:
llm_answer_raw = llm_utils.invoke_claude(
prompt, "sonnet_latest", filename, usage_label="SPECIAL_CASE_BREAKOUT"
)
llm_answer_final = _parser(llm_answer_raw)
return llm_answer_final
def prompt_lob_relationship(
answer_dict: dict, field: str, exhibit_text: str, filename: str
):
@@ -44,107 +44,6 @@ def prompt_exhibit_level(
return llm_answer_final
def prompt_exhibit_level_breakout(
exhibit_level_answers: dict[str, str], filename: str
) -> dict[str, str]:
"""
Process exhibit-level answers by breaking out facility adjustment terms using an LLM.
This function checks if a FACILITY_ADJUSTMENT_TERM exists in the exhibit-level answers,
and if present, invokes Claude LLM to break it down into structured components. The
resulting parsed data is then merged back into the exhibit-level answers dictionary.
Args:
exhibit_level_answers (dict[str, str]): A dictionary containing exhibit-level
field names as keys and their corresponding values as strings. Must contain
a "FACILITY_ADJUSTMENT_TERM" key to trigger processing.
filename (str): The name of the file being processed, used for logging and
tracking purposes in the LLM invocation.
Returns:
dict[str, str]: The updated exhibit_level_answers dictionary with additional
fields extracted from the FACILITY_ADJUSTMENT_TERM breakout, or the original
dictionary if no FACILITY_ADJUSTMENT_TERM was present.
Raises:
Any exceptions raised by llm_utils.invoke_claude() or _parser()
will propagate to the caller.
"""
# FACILITY_ADJUSTMENT_BREAKOUT
if not string_utils.is_empty(
exhibit_level_answers.get("FACILITY_ADJUSTMENT_TERM", "")
):
prompt, _parser = prompt_templates.FACILITY_ADJUSTMENT_BREAKOUT(
exhibit_level_answers["FACILITY_ADJUSTMENT_TERM"]
)
llm_answer_raw = llm_utils.invoke_claude(
prompt=prompt, model_id="sonnet_latest", filename=filename
)
llm_answer_final = _parser(llm_answer_raw)
exhibit_level_answers.update(llm_answer_final)
return exhibit_level_answers
def prompt_dynamic_primary(
exhibit_text: str, field: Field, constants: Constants, filename: str, TEMPLATE
):
# Check if we should use context caching for DYNAMIC_PRIMARY
if TEMPLATE == prompt_templates.DYNAMIC_PRIMARY:
# Use the context caching version that splits context from field question
context_text, prompt, _parser = prompt_templates.DYNAMIC_PRIMARY(
exhibit_text,
field.field_name,
field.get_prompt(constants),
)
logging.debug(
f"Dynamic primary prompt with context caching for {filename}; {field}: {prompt}"
)
logging.debug(f"Context length for caching: {len(context_text)} chars")
# Invoke with context_for_caching to enable exhibit-level caching
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.DYNAMIC_PRIMARY_INSTRUCTION(),
context_for_caching=context_text,
usage_label="DYNAMIC_PRIMARY",
)
else:
# Support both legacy (prompt, parser) and context-aware
# (context_text, prompt, parser) template return signatures.
template_result = TEMPLATE(
exhibit_text, field.field_name, field.get_prompt(constants)
)
if not isinstance(template_result, (tuple, list)):
raise ValueError(
f"Template must return tuple or list, got {type(template_result)} "
f"for {getattr(TEMPLATE, '__name__', TEMPLATE)}"
)
if len(template_result) == 3:
context_text, prompt, _parser = template_result
elif len(template_result) == 2:
prompt, _parser = template_result
context_text = None
else:
raise ValueError(
f"Unexpected template return size for {getattr(TEMPLATE, '__name__', TEMPLATE)}: {len(template_result)}"
)
logging.debug(f"Dynamic primary prompt for {filename}; {field}: {prompt}")
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.DYNAMIC_PRIMARY_INSTRUCTION(),
context_for_caching=context_text,
usage_label="DYNAMIC_PRIMARY",
)
logging.debug(f"Claude answer for {filename}; {field}: {llm_answer_raw}")
llm_answer_final = _parser(llm_answer_raw)
return llm_answer_final
def prompt_reimbursement_primary(
page_text: str,
filename: str,
@@ -295,37 +194,6 @@ def prompt_fee_schedule_breakout(
return fs_breakout_dict
def prompt_grouper_breakout(
service: str,
reimbursement_method: str,
filename: str,
):
"""
Call GROUPER_BREAKOUT prompt with cached instruction.
Note: Field definitions are now included in GROUPER_BREAKOUT_INSTRUCTION() for caching.
"""
prompt, _parser = prompt_templates.GROUPER_BREAKOUT(
service,
reimbursement_method,
)
logging.debug(f"Grouper Breakout Prompt for {filename}: {prompt}")
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=prompt_templates.GROUPER_BREAKOUT_INSTRUCTION(),
usage_label="GROUPER_BREAKOUT",
)
logging.debug(f"LLM Response for {filename}: {llm_answer_raw}")
try:
grouper_breakout_dict = _parser(llm_answer_raw)
except:
grouper_breakout_dict = {}
return grouper_breakout_dict
def prompt_carveout_check(
service_term: str,
reimb_term: str,
@@ -352,39 +220,6 @@ def prompt_carveout_check(
return carveout_answer
def prompt_special_case_breakout(breakout_template, term_answer, filename):
# Get the prompt from the template
prompt, _parser = breakout_template(term_answer)
# Ensure prompt is a string for Bedrock messages API
if not isinstance(prompt, str):
logging.warning(
f"Expected prompt string, got {type(prompt)} in prompt_special_case_breakout; coercing to str."
)
prompt = str(prompt)
# Look for a corresponding _INSTRUCTION() function for caching
template_name = getattr(breakout_template, "__name__", "SPECIAL_CASE_BREAKOUT")
instruction_func_name = f"{template_name}_INSTRUCTION"
instruction_func = getattr(prompt_templates, instruction_func_name, None)
if instruction_func and callable(instruction_func):
llm_answer_raw = llm_utils.invoke_claude(
prompt,
"sonnet_latest",
filename,
cache=True,
instruction=instruction_func(),
usage_label=template_name,
)
else:
llm_answer_raw = llm_utils.invoke_claude(
prompt, "sonnet_latest", filename, usage_label="SPECIAL_CASE_BREAKOUT"
)
llm_answer_final = _parser(llm_answer_raw)
return llm_answer_final
def prompt_lob_relationship(
answer_dict: dict, field: str, exhibit_text: str, filename: str
):
+9 -8
View File
@@ -11,7 +11,6 @@ Usage:
"""
import concurrent.futures
import importlib
import logging
import os
import traceback
@@ -28,6 +27,7 @@ import src.utils.io_utils as io_utils
import src.utils.llm_utils as llm_utils
import src.utils.logging_utils as logging_utils
import src.utils.usage_tracking as usage_tracking
from src.pipelines.runtime_context import set_active_client
from src.constants.constants import Constants
from src import config
from src.core.registry import get_registry
@@ -53,19 +53,17 @@ def get_file_processing_module(client: str):
Returns:
An object with a process_file(file_object, constants, run_timestamp) method.
"""
if client == "saas":
from src.pipelines.saas import file_processing
return file_processing
registry = get_registry()
if registry.is_vendor(client):
from src.pipelines.vendors.shared.generic_processor import VendorProcessor
return VendorProcessor(client)
module_path = f"src.pipelines.clients.{client}.file_processing"
return importlib.import_module(module_path)
# Set active client so shared resolver can select function-level overrides.
set_active_client(client)
from src.pipelines.shared import file_processing
return file_processing
def safe_process_file(item, constants, run_timestamp, file_processing):
@@ -141,6 +139,9 @@ def main(client: str = "saas", testing=False, test_params={}):
f"Unknown client/vendor: {client}. " f"Available: {registry.list_clients()}"
)
# Ensure dynamic fallback resolution targets the requested client.
set_active_client(client)
# Get client-specific file processing module
file_processing = get_file_processing_module(client)
logging.info(f"Using {client} pipeline")
+31
View File
@@ -0,0 +1,31 @@
"""Runtime pipeline context for ClientResolver — stores the active client
so the ClientResolver modules can dispatch to client-specific or SaaS functions."""
_active_client = "saas"
def set_active_client(client_name: str) -> None:
"""Set the active client for dynamic module resolution."""
global _active_client
normalized = (client_name or "saas").lower()
_active_client = normalized
def get_active_client() -> str:
"""Get the active client, defaulting to config or saas."""
if _active_client:
return _active_client
try:
from src import config
configured = getattr(config, "CLIENT", ["saas"])
if isinstance(configured, list) and configured:
return str(configured[0]).lower()
if configured:
return str(configured).lower()
except Exception:
# Keep resolver robust even if config import is unavailable.
pass
return "saas"
+3
View File
@@ -11,6 +11,7 @@ import pandas as pd
random.seed(42)
import src.pipelines.saas.file_processing as file_processing
from src.pipelines.runtime_context import set_active_client
from src.pipelines.shared.extraction import one_to_one_funcs
from src.pipelines.shared.postprocessing import postprocessing_funcs
from src.constants.investment_columns import FIELD_FORMAT_MAPPING
@@ -94,6 +95,8 @@ def safe_process_file(item, constants, run_timestamp):
def main(testing=False, test_params={}):
set_active_client("saas")
# Check if batch_id is specified
if not config.BATCH_ID:
raise ValueError(
@@ -2,7 +2,7 @@ import logging
import concurrent.futures
from typing import TYPE_CHECKING, Optional
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src.pipelines.shared.preprocessing import preprocessing_funcs
from src.utils import string_utils, timing_utils
from src.constants.constants import Constants
@@ -34,7 +34,7 @@ from src.pipelines.shared.preprocessing.exhibit_smart_chunking_funcs import (
if TYPE_CHECKING:
from src.pipelines.shared.extraction.page_funcs import Page
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
import logging
@@ -2,7 +2,7 @@ import concurrent.futures
import logging
import pandas as pd
from src.pipelines.shared.extraction import dynamic_funcs
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src.pipelines.shared.postprocessing import aarete_derived, postprocessing_funcs
from src.utils import llm_utils, string_utils, timing_utils
import src.prompts.prompt_templates as prompt_templates
@@ -9,7 +9,7 @@ from collections import defaultdict
import pandas as pd
import src.pipelines.shared.postprocessing.postprocessing_funcs as postprocessing_funcs
import src.pipelines.saas.prompts.prompt_calls as prompt_calls
import src.pipelines.shared.prompts.prompt_calls as prompt_calls
import src.utils.string_utils as string_utils
from src.constants.constants import Constants
from src import config
+1 -1
View File
@@ -1,7 +1,7 @@
import logging
from copy import deepcopy
import pandas as pd
import src.pipelines.saas.prompts.prompt_calls as prompt_calls
import src.pipelines.shared.prompts.prompt_calls as prompt_calls
from src.constants.constants import Constants
from src.pipelines.shared.postprocessing import aarete_derived
from src.pipelines.shared.extraction import one_to_one_funcs
@@ -7,7 +7,7 @@ import src.constants.regex_patterns as regex_patterns
import src.config as config
import src.prompts.prompt_templates as prompt_templates
from src.utils import llm_utils, string_utils, json_utils
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src.prompts.fieldset import Field, FieldSet
import pandas as pd
+51
View File
@@ -0,0 +1,51 @@
"""ClientResolver — file processing dispatch.
Resolves file processing functions at runtime:
1. If the active client defines the function, the client implementation is used.
2. If the client does not define the function, the SaaS implementation is reused.
Client modules only need to implement functions that differ from SaaS.
"""
import importlib
from src.pipelines.runtime_context import get_active_client
_SAAS_MODULE = "src.pipelines.saas.file_processing"
class ClientResolver:
"""Resolves file processing functions between client-specific and SaaS implementations."""
@staticmethod
def get_client_module_name(client_name: str) -> str:
return f"src.pipelines.clients.{client_name}.file_processing"
@staticmethod
def load_module(module_name: str):
try:
return importlib.import_module(module_name)
except ImportError as exc:
if exc.name == module_name:
return None
raise
@classmethod
def resolve(cls, name: str):
saas_module = importlib.import_module(_SAAS_MODULE)
client_name = get_active_client()
if client_name != "saas":
client_module = cls.load_module(cls.get_client_module_name(client_name))
if client_module is not None and hasattr(client_module, name):
return getattr(client_module, name)
return getattr(saas_module, name)
_resolver = ClientResolver()
def __getattr__(name: str):
return _resolver.resolve(name)
@@ -4,7 +4,7 @@ from collections import Counter, defaultdict
from typing import TYPE_CHECKING
from src import config
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src.pipelines.shared.preprocessing import preprocessing_funcs
from src.utils import timing_utils
@@ -3,7 +3,7 @@ import re
from typing import TYPE_CHECKING, Optional
import src.utils.string_utils as string_utils
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src import config
from src.constants.delimiters import TABLE_START_MARKER, TABLE_END_MARKER
+1
View File
@@ -0,0 +1 @@
"""Shared prompt module proxies with client fallback."""
@@ -0,0 +1,51 @@
"""ClientResolver — prompt call dispatch.
Resolves prompt functions at runtime:
1. If the active client defines the function, the client implementation is used.
2. If the client does not define the function, the SaaS implementation is reused.
Client modules only need to implement functions that differ from SaaS.
"""
import importlib
from src.pipelines.runtime_context import get_active_client
_SAAS_MODULE = "src.pipelines.saas.prompts.prompt_calls"
class ClientResolver:
"""Resolves prompt call functions between client-specific and SaaS implementations."""
@staticmethod
def get_client_module_name(client_name: str) -> str:
return f"src.pipelines.clients.{client_name}.prompts.prompt_calls"
@staticmethod
def load_module(module_name: str):
try:
return importlib.import_module(module_name)
except ImportError as exc:
if exc.name == module_name:
return None
raise
@classmethod
def resolve(cls, name: str):
saas_module = importlib.import_module(_SAAS_MODULE)
client_name = get_active_client()
if client_name != "saas":
client_module = cls.load_module(cls.get_client_module_name(client_name))
if client_module is not None and hasattr(client_module, name):
return getattr(client_module, name)
return getattr(saas_module, name)
_resolver = ClientResolver()
def __getattr__(name: str):
return _resolver.resolve(name)
+7 -4
View File
@@ -6,7 +6,8 @@ from unittest.mock import MagicMock, patch
import pandas as pd
from src.constants.constants import Constants
from src.pipelines.shared import dynamic_funcs
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.shared.prompts import prompt_calls
from src.pipelines.runtime_context import set_active_client
from src.prompts.fieldset import Field, FieldSet
@@ -15,6 +16,8 @@ class TestDynamicFuncsWithRealConstants(unittest.TestCase):
def setUp(self):
"""Set up test fixtures before each test method."""
set_active_client("saas")
# Create real Constants object
self.constants = Constants()
@@ -221,7 +224,7 @@ class TestDynamicFuncsWithRealConstants(unittest.TestCase):
self.answer_dicts, self.constants
)
@patch("src.pipelines.saas.prompts.prompt_calls.prompt_dynamic")
@patch("src.pipelines.shared.extraction.dynamic_funcs.prompt_calls.prompt_dynamic")
def test_dynamic_with_real_constants(self, mock_prompt_dynamic):
"""Test dynamic function with real Constants object."""
# Setup
@@ -270,7 +273,7 @@ class TestDynamicFuncsWithRealConstants(unittest.TestCase):
# Verify that prompt_dynamic was called twice
self.assertEqual(mock_prompt_dynamic.call_count, 2)
@patch("src.pipelines.saas.prompts.prompt_calls.prompt_dynamic")
@patch("src.pipelines.shared.extraction.dynamic_funcs.prompt_calls.prompt_dynamic")
def test_dynamic_with_no_fields(self, mock_prompt_dynamic):
"""Test dynamic function with an empty fieldset."""
# Setup
@@ -294,7 +297,7 @@ class TestDynamicFuncsWithRealConstants(unittest.TestCase):
self.assertEqual(len(updated_reimbursement_fields.fields), 0)
mock_prompt_dynamic.assert_not_called()
@patch("src.pipelines.saas.prompts.prompt_calls.prompt_dynamic")
@patch("src.pipelines.shared.extraction.dynamic_funcs.prompt_calls.prompt_dynamic")
def test_dynamic_with_no_answers(self, mock_prompt_dynamic):
"""Test dynamic function when neither header nor text contains answers."""
# Setup
+110
View File
@@ -0,0 +1,110 @@
import types
import unittest
from unittest.mock import patch
import src.pipelines.shared.file_processing as shared_file_processing
from src.pipelines.runtime_context import set_active_client
class TestFileProcessingResolver(unittest.TestCase):
"""Contract tests for shared file processing resolver behavior."""
def setUp(self):
set_active_client("saas")
def tearDown(self):
set_active_client("saas")
@staticmethod
def _make_module(module_name: str, function_name: str, function_obj):
module = types.ModuleType(module_name)
setattr(module, function_name, function_obj)
return module
def test_file_processing_resolver_uses_client_override_when_present(self):
def saas_fn():
return "saas"
def client_fn():
return "client"
function_name = "target_file_fn"
saas_module = self._make_module("saas_file_processing", function_name, saas_fn)
client_module = self._make_module(
"client_file_processing", function_name, client_fn
)
set_active_client("test_client")
with (
patch.object(
shared_file_processing.importlib,
"import_module",
return_value=saas_module,
),
patch.object(
shared_file_processing.ClientResolver,
"load_module",
return_value=client_module,
),
):
resolved = shared_file_processing.__getattr__(function_name)
self.assertIs(resolved, client_fn)
self.assertEqual(resolved(), "client")
def test_file_processing_resolver_falls_back_to_saas_when_function_missing(self):
def saas_fn():
return "saas"
function_name = "target_file_fn"
saas_module = self._make_module("saas_file_processing", function_name, saas_fn)
client_module = types.ModuleType("client_file_processing")
set_active_client("test_client")
with (
patch.object(
shared_file_processing.importlib,
"import_module",
return_value=saas_module,
),
patch.object(
shared_file_processing.ClientResolver,
"load_module",
return_value=client_module,
),
):
resolved = shared_file_processing.__getattr__(function_name)
self.assertIs(resolved, saas_fn)
self.assertEqual(resolved(), "saas")
def test_file_processing_resolver_falls_back_to_saas_when_client_module_missing(
self,
):
def saas_fn():
return "saas"
function_name = "target_file_fn"
saas_module = self._make_module("saas_file_processing", function_name, saas_fn)
set_active_client("missing_client")
with (
patch.object(
shared_file_processing.importlib,
"import_module",
return_value=saas_module,
),
patch.object(
shared_file_processing.ClientResolver,
"load_module",
return_value=None,
),
):
resolved = shared_file_processing.__getattr__(function_name)
self.assertIs(resolved, saas_fn)
self.assertEqual(resolved(), "saas")
if __name__ == "__main__":
unittest.main()
+3 -1
View File
@@ -1,7 +1,7 @@
import unittest
from unittest.mock import Mock, patch, MagicMock
from src.pipelines.shared import one_to_n_funcs
from src.pipelines.saas.prompts import prompt_calls
from src.pipelines.runtime_context import set_active_client
from src.prompts.fieldset import FieldSet
from src.constants.constants import Constants
import json
@@ -12,6 +12,8 @@ class TestOneToNFuncs(unittest.TestCase):
def setUp(self):
"""Set up test fixtures."""
set_active_client("saas")
self.constants = Constants()
self.filename = "test_contract.pdf"
self.exhibit_text = "This is a sample exhibit text with reimbursement terms."
-1
View File
@@ -16,7 +16,6 @@ import pytest
from src.pipelines.shared.extraction.exhibit_funcs import Exhibit, get_exhibit_list
from src.pipelines.shared.extraction.page_funcs import Page
from src.pipelines.shared import preprocessing_funcs
from src.pipelines.saas import file_processing
from src.pipelines.shared import dynamic_funcs
from src.prompts.fieldset import FieldSet
+100
View File
@@ -0,0 +1,100 @@
import types
import unittest
from unittest.mock import patch
import src.pipelines.shared.prompts.prompt_calls as shared_prompt_calls
from src.pipelines.runtime_context import set_active_client
class TestPromptResolver(unittest.TestCase):
"""Contract tests for shared prompt resolver behavior."""
def setUp(self):
set_active_client("saas")
def tearDown(self):
set_active_client("saas")
@staticmethod
def _make_module(module_name: str, function_name: str, function_obj):
module = types.ModuleType(module_name)
setattr(module, function_name, function_obj)
return module
def test_prompt_resolver_uses_client_override_when_present(self):
def saas_fn():
return "saas"
def client_fn():
return "client"
function_name = "target_prompt_fn"
saas_module = self._make_module("saas_prompts", function_name, saas_fn)
client_module = self._make_module("client_prompts", function_name, client_fn)
set_active_client("test_client")
with (
patch.object(
shared_prompt_calls.importlib, "import_module", return_value=saas_module
),
patch.object(
shared_prompt_calls.ClientResolver,
"load_module",
return_value=client_module,
),
):
resolved = shared_prompt_calls.__getattr__(function_name)
self.assertIs(resolved, client_fn)
self.assertEqual(resolved(), "client")
def test_prompt_resolver_falls_back_to_saas_when_function_missing(self):
def saas_fn():
return "saas"
function_name = "target_prompt_fn"
saas_module = self._make_module("saas_prompts", function_name, saas_fn)
client_module = types.ModuleType("client_prompts")
set_active_client("test_client")
with (
patch.object(
shared_prompt_calls.importlib, "import_module", return_value=saas_module
),
patch.object(
shared_prompt_calls.ClientResolver,
"load_module",
return_value=client_module,
),
):
resolved = shared_prompt_calls.__getattr__(function_name)
self.assertIs(resolved, saas_fn)
self.assertEqual(resolved(), "saas")
def test_prompt_resolver_falls_back_to_saas_when_client_module_missing(self):
def saas_fn():
return "saas"
function_name = "target_prompt_fn"
saas_module = self._make_module("saas_prompts", function_name, saas_fn)
set_active_client("missing_client")
with (
patch.object(
shared_prompt_calls.importlib, "import_module", return_value=saas_module
),
patch.object(
shared_prompt_calls.ClientResolver,
"load_module",
return_value=None,
),
):
resolved = shared_prompt_calls.__getattr__(function_name)
self.assertIs(resolved, saas_fn)
self.assertEqual(resolved(), "saas")
if __name__ == "__main__":
unittest.main()