228 lines
7.8 KiB
Python
228 lines
7.8 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 datetime import datetime
|
|
from typing import Any, cast
|
|
from unittest.mock import patch
|
|
|
|
from cleveragents.application.services.plan_service import PlanService
|
|
from cleveragents.domain.models.core import (
|
|
Actor,
|
|
Change,
|
|
OperationType,
|
|
Plan,
|
|
Project,
|
|
)
|
|
from cleveragents.domain.models.core import (
|
|
Context as PlanContext,
|
|
)
|
|
from cleveragents.domain.providers.ai_provider import (
|
|
ActorInvocationContext,
|
|
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",
|
|
)
|
|
config_blob: dict[str, Any] = field(
|
|
default_factory=lambda: cast(dict[str, Any], {})
|
|
)
|
|
generate_calls: int = field(init=False, default=0)
|
|
stream_calls: int = field(init=False, default=0)
|
|
last_actor_context: ActorInvocationContext | None = field(init=False, default=None)
|
|
_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],
|
|
actor_context: ActorInvocationContext | None = None,
|
|
progress_callback: Any | None = None,
|
|
) -> ProviderResponse:
|
|
self.last_actor_context = actor_context
|
|
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],
|
|
actor_context: ActorInvocationContext | None = None,
|
|
progress_callback: Any | None = None,
|
|
) -> Iterator[dict[str, Any]]:
|
|
self.last_actor_context = actor_context
|
|
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 actor-based resolver 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,
|
|
actor_name: str | None,
|
|
*,
|
|
provider: str | None = None,
|
|
model: str | None = None,
|
|
) -> tuple[Any, str, str, Actor, dict[str, Any], ActorInvocationContext]:
|
|
requested = (provider or actor_name or default_key).lower()
|
|
stub = normalized.get(requested, normalized[default_key])
|
|
context.provider_resolver_calls.append( # type: ignore[attr-defined]
|
|
{"requested": actor_name or provider, "resolved": stub.name}
|
|
)
|
|
now = datetime.now()
|
|
config_blob: dict[str, Any] = dict(cast(dict[str, Any], stub.config_blob))
|
|
graph_descriptor = config_blob.get("graph_descriptor")
|
|
graph_descriptor_typed: dict[str, Any] | None = (
|
|
cast(dict[str, Any], graph_descriptor)
|
|
if isinstance(graph_descriptor, dict)
|
|
else None
|
|
)
|
|
initial_context_raw = config_blob.get("initial_context") or config_blob.get(
|
|
"context_variables"
|
|
)
|
|
initial_context: dict[str, Any] = (
|
|
dict(initial_context_raw)
|
|
if isinstance(initial_context_raw, dict)
|
|
else {}
|
|
)
|
|
options_raw = config_blob.get("options")
|
|
options_payload: dict[str, Any] = (
|
|
dict(options_raw) if isinstance(options_raw, dict) else {}
|
|
)
|
|
actor = Actor(
|
|
id=None,
|
|
name=f"local/{stub.name}",
|
|
provider=stub.name,
|
|
model=stub.model_id or f"{stub.name}-model",
|
|
config_blob=config_blob,
|
|
config_hash=Actor.compute_hash(config_blob),
|
|
graph_descriptor=graph_descriptor_typed,
|
|
unsafe=bool(config_blob.get("unsafe", True)),
|
|
is_built_in=True,
|
|
is_default=True,
|
|
created_at=now,
|
|
updated_at=now,
|
|
)
|
|
actor_context = ActorInvocationContext(
|
|
name=actor.name,
|
|
provider=actor.provider,
|
|
model=actor.model,
|
|
options=options_payload,
|
|
graph_descriptor=graph_descriptor_typed,
|
|
unsafe=actor.unsafe,
|
|
config_hash=actor.config_hash,
|
|
config_blob=config_blob,
|
|
initial_context=initial_context,
|
|
)
|
|
return (
|
|
stub,
|
|
stub.name,
|
|
stub.model_id or f"{stub.name}-model",
|
|
actor,
|
|
{
|
|
"actor": actor.name,
|
|
"actor_source": "stub",
|
|
"actor_is_default": actor.is_default,
|
|
"actor_unsafe": actor.unsafe,
|
|
"default_provider": stub.name,
|
|
"default_provider_source": "stub",
|
|
"default_model": stub.model_id or f"{stub.name}-model",
|
|
"default_model_source": "stub",
|
|
},
|
|
actor_context,
|
|
)
|
|
|
|
patcher = patch.object(PlanService, "_resolve_ai_provider_for_actor", _resolver)
|
|
patcher.start()
|
|
cleanup = getattr(context, "add_cleanup", None)
|
|
if callable(cleanup):
|
|
cleanup(patcher.stop)
|