Spaces:
Running
Running
| import os | |
| from typing import Optional | |
| from .registry.persona import VISIT_TYPE | |
| from .utils import colorstr, log | |
| from .client import GeminiClient, GeminiVertexClient, GPTClient, GPTAzureClient | |
| _PROMPT_DIR = os.path.join(os.path.dirname(__file__), "assets", "prompt") | |
| class CheckerAgent: | |
| def __init__(self, | |
| model: str, | |
| visit_type: str = 'emergency_department', | |
| api_key: Optional[str] = None, | |
| use_azure: bool = False, | |
| use_vertex: bool = False, | |
| azure_endpoint: Optional[str] = None, | |
| user_prompt_path: Optional[str] = None, | |
| **kwargs) -> None: | |
| self.visit_type = visit_type.lower() | |
| self.__sanity_check() | |
| self.model = model | |
| self.random_seed = kwargs.get('random_seed', None) | |
| self.temperature = kwargs.get('temperature', 0.0) | |
| self._init_model( | |
| model=self.model, | |
| api_key=api_key, | |
| use_azure=use_azure, | |
| use_vertex=use_vertex, | |
| azure_endpoint=azure_endpoint, | |
| ) | |
| self.prompt_template = self._init_prompt(self.visit_type, user_prompt_path) | |
| log("CheckerAgent initialized successfully", color=True) | |
| def _init_model(self, | |
| model: str, | |
| api_key: Optional[str] = None, | |
| use_azure: bool = False, | |
| use_vertex: bool = False, | |
| azure_endpoint: Optional[str] = None) -> None: | |
| if 'gemini' in self.model.lower(): | |
| self.client = GeminiVertexClient(model, api_key) if use_vertex else GeminiClient(model, api_key) | |
| elif 'gpt' in self.model.lower(): | |
| self.client = GPTAzureClient(model, api_key, azure_endpoint) if use_azure else GPTClient(model, api_key) | |
| else: | |
| raise ValueError(colorstr("red", f"Unsupported model: {self.model}. Supported models are 'gemini' and 'gpt'.")) | |
| def _init_prompt(self, visit_type: str, user_prompt_path: Optional[str] = None) -> str: | |
| if not user_prompt_path: | |
| fname = "op_terminate_user.txt" if visit_type == 'outpatient' else "ed_terminate_user.txt" | |
| with open(os.path.join(_PROMPT_DIR, fname), 'r') as f: | |
| return f.read() | |
| else: | |
| if not os.path.exists(user_prompt_path): | |
| raise FileNotFoundError(colorstr("red", f"User prompt file not found: {user_prompt_path}")) | |
| with open(user_prompt_path, 'r') as f: | |
| return f.read() | |
| def reset_history(self, verbose: bool = True) -> None: | |
| self.client.reset_history(verbose=verbose) | |
| def __sanity_check(self) -> None: | |
| if self.visit_type not in VISIT_TYPE: | |
| raise ValueError(colorstr("red", f"Invalid visiting type: {self.visit_type}. Supported types: {', '.join(VISIT_TYPE)}")) | |
| def __call__(self, response: str, **kwargs) -> str: | |
| # Escape braces in `response` before substitution to prevent unintended | |
| # format-string expansion if the LLM output contains `{...}` patterns. | |
| safe_response = response.replace("{", "{{").replace("}", "}}") | |
| return self.client( | |
| user_prompt=self.prompt_template.format(response=safe_response), | |
| using_multi_turn=False, | |
| verbose=False, | |
| temperature=self.temperature, | |
| seed=self.random_seed, | |
| **kwargs | |
| ) | |