"""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)