igerasimov's picture
MVP Milestone 5
d840c10
Raw
History Blame Contribute Delete
2.03 kB
"""Deterministic fake model client for unit tests."""
from __future__ import annotations
from collections.abc import Iterable
from typing import TypeAlias
from pydantic import BaseModel, ValidationError
from gcmd_classifier.errors import NonRetryableModelError, StructuredModelResponseError
from gcmd_classifier.llm.base import ModelRequest, ModelResponse, StructuredResponseT
ScriptedAction: TypeAlias = BaseModel | dict | Exception
class FakeModelClient:
"""Scripted model client that never imports or calls a live provider."""
def __init__(self, actions: Iterable[ScriptedAction]) -> None:
self._actions = list(actions)
self.requests: list[ModelRequest] = []
def generate_structured(
self,
request: ModelRequest[StructuredResponseT],
) -> ModelResponse[StructuredResponseT]:
"""Return the next scripted response or raise the next scripted exception."""
self.requests.append(request)
if not self._actions:
raise NonRetryableModelError("Fake model script is exhausted.")
action = self._actions.pop(0)
if isinstance(action, Exception):
raise action
try:
parsed = _parse_action(action, request.response_schema)
except ValidationError as exc:
raise StructuredModelResponseError(
"Fake model response failed schema validation."
) from exc
return ModelResponse(
parsed=parsed,
provider=request.provider,
model_name=request.model_name,
prompt_version=request.prompt_version,
)
@property
def remaining_actions(self) -> int:
"""Number of scripted actions that have not been consumed."""
return len(self._actions)
def _parse_action(
action: BaseModel | dict,
response_schema: type[StructuredResponseT],
) -> StructuredResponseT:
if isinstance(action, response_schema):
return action
return response_schema.model_validate(action)