Files
cleveragents-core/features/steps/provider_stub_utils.py
T

152 lines
4.7 KiB
Python

"""Shared provider stub utilities for Behave tests.
These helpers let streaming and auto-debug scenarios stub provider selection
logic without touching production code.
"""
from __future__ import annotations
from collections.abc import Iterator
from dataclasses import dataclass, field
from typing import Any
from unittest.mock import patch
from cleveragents.application.services.plan_service import PlanService
from cleveragents.domain.models.core import (
Change,
OperationType,
Plan,
Project,
)
from cleveragents.domain.models.core import (
Context as PlanContext,
)
from cleveragents.domain.providers.ai_provider import ProviderResponse
@dataclass
class RecordingProviderStub:
"""Deterministic provider stub that records usage."""
name: str
token_count: int = 123
model_id: str | None = None
should_fail: bool = False
failures_before_success: int = 0
stream_nodes: tuple[str, ...] = (
"load_context",
"analyze_requirements",
"generate_plan",
"validate",
)
generate_calls: int = field(init=False, default=0)
stream_calls: int = field(init=False, default=0)
_failure_count: int = field(init=False, default=0)
def __post_init__(self) -> None:
if not self.model_id:
self.model_id = f"{self.name}-model"
def _build_change(self, plan: Plan) -> Change:
plan_id = plan.id or 1
return Change(
id=None,
plan_id=plan_id,
file_path=f"{self.name}_generated.py",
operation=OperationType.CREATE,
original_content=None,
new_content=f"# generated by {self.name}\nprint('hello from {self.name}')\n",
applied=False,
applied_at=None,
new_path=None,
)
def _build_response(self, plan: Plan) -> ProviderResponse:
return ProviderResponse(
changes=[self._build_change(plan)],
model_used=self.model_id or f"{self.name}-model",
token_count=self.token_count,
error_message=None,
)
def _should_raise(self) -> bool:
if self.should_fail:
return True
if self.failures_before_success <= 0:
return False
if self._failure_count < self.failures_before_success:
self._failure_count += 1
return True
return False
def generate_changes(
self,
project: Project,
plan: Plan,
contexts: list[PlanContext],
progress_callback: Any | None = None,
) -> ProviderResponse:
self.generate_calls += 1
if self._should_raise():
raise Exception(f"{self.name} provider failure")
if progress_callback:
progress_callback(100)
return self._build_response(plan)
def stream_changes(
self,
project: Project,
plan: Plan,
contexts: list[PlanContext],
progress_callback: Any | None = None,
) -> Iterator[dict[str, Any]]:
self.stream_calls += 1
if self._should_raise():
raise Exception(f"{self.name} provider failure")
response = self._build_response(plan)
for node in self.stream_nodes:
payload: dict[str, Any] = {
node: {"status": "running" if node != "generate_plan" else "completed"}
}
yield payload
yield {"__end__": {"response": response}}
def install_provider_resolver_patch(
context: Any,
*,
override_map: dict[str, RecordingProviderStub],
default_provider: str,
) -> None:
"""Patch PlanService._resolve_ai_provider to return test stubs."""
normalized = {key.lower(): stub for key, stub in override_map.items()}
default_key = default_provider.lower()
if default_key not in normalized:
raise ValueError("Default provider must exist in override_map")
call_list = getattr(context, "provider_resolver_calls", None)
if isinstance(call_list, list):
call_list.clear()
else:
call_list = []
context.provider_resolver_calls = call_list
def _resolver(
self: PlanService,
provider: str | None,
model: str | None,
) -> tuple[Any, str, str]:
requested = (provider or default_key).lower()
stub = normalized.get(requested, normalized[default_key])
context.provider_resolver_calls.append( # type: ignore[attr-defined]
{"requested": provider, "resolved": stub.name}
)
return stub, stub.name, stub.model_id or f"{stub.name}-model"
patcher = patch.object(PlanService, "_resolve_ai_provider", _resolver)
patcher.start()
cleanup = getattr(context, "add_cleanup", None)
if callable(cleanup):
cleanup(patcher.stop)