152 lines
4.7 KiB
Python
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)
|