from __future__ import annotations import ast from typing import Any from unittest.mock import MagicMock, patch from behave import given, then, when from cleveragents.providers.llm.google_provider import GoogleChatProvider def _register_cleanup(context, cleanup): if hasattr(context, "add_cleanup"): context.add_cleanup(cleanup) else: cleanup_handlers = getattr(context, "_cleanup_handlers", []) cleanup_handlers.append(cleanup) context._cleanup_handlers = cleanup_handlers def _parse_kwargs_string(kwargs_string: str) -> dict[str, Any]: if not kwargs_string: return {} parsed: dict[str, Any] = {} for entry in kwargs_string.split(","): entry = entry.strip() if not entry: continue if "=" not in entry: parsed[entry] = True continue key, value = entry.split("=", 1) key = key.strip() value = value.strip() try: parsed_value = ast.literal_eval(value) except Exception: parsed_value = value parsed[key] = parsed_value return parsed def _setup_plan_generation_graph(context) -> MagicMock: patcher = patch( "cleveragents.providers.llm.langchain_chat_provider.PlanGenerationGraph" ) context.plan_generation_patcher = patcher mock_graph_class = patcher.start() _register_cleanup(context, patcher.stop) mock_graph_instance = MagicMock(name="PlanGenerationGraphInstance") default_state = { "generated_changes": [], "validation_result": {"status": "PASS"}, "error": None, } state_override = getattr(context, "plan_generation_state_override", None) mock_graph_instance.invoke.return_value = state_override or default_state invoke_side_effect = getattr(context, "plan_generation_invoke_side_effect", None) if invoke_side_effect is not None: mock_graph_instance.invoke.side_effect = invoke_side_effect stream_events = getattr(context, "plan_generation_stream_events", None) if stream_events is None: mock_graph_instance.stream.return_value = iter(()) else: events_copy = list(stream_events) def _stream(*_args, **_kwargs): for event in events_copy: yield event mock_graph_instance.stream.side_effect = _stream mock_graph_class.return_value = mock_graph_instance context.plan_generation_graph = mock_graph_instance context.plan_generation_graph_class = mock_graph_class return mock_graph_instance @given('I set the Google provider extra kwargs "{kwargs_string}"') def step_set_google_provider_kwargs(context, kwargs_string): context.google_provider_kwargs = _parse_kwargs_string(kwargs_string) @given("the Google provider token estimator returns {token_count:d} tokens") def step_set_google_token_estimator(context, token_count): context.google_token_count = token_count @when('I create a Google chat provider with API key "{api_key}" and model "{model_id}"') def step_create_google_provider(context, api_key, model_id): patcher = patch("cleveragents.providers.llm.google_provider.ChatGoogleGenerativeAI") context.chat_google_patcher = patcher mock_chat_class = patcher.start() _register_cleanup(context, patcher.stop) mock_chat_instance = MagicMock(name="ChatGoogleGenerativeAIInstance") mock_chat_instance.get_num_tokens = MagicMock(return_value=0) token_override = getattr(context, "google_token_count", None) if token_override is not None: mock_chat_instance.get_num_tokens.return_value = token_override mock_chat_class.return_value = mock_chat_instance extra_kwargs = dict(getattr(context, "google_provider_kwargs", {})) context.provider = GoogleChatProvider( api_key=api_key, model=model_id, **extra_kwargs, ) context.chat_google_class = mock_chat_class context.chat_google_instance = mock_chat_instance @when("I attempt to create a Google chat provider without an API key") def step_google_provider_without_api_key(context): context.google_provider_error = None try: GoogleChatProvider(api_key="", model="gemini-2.0-flash") except Exception as exc: # pragma: no cover - defensive context.google_provider_error = exc @when("I request plan generation from the Google provider") def step_request_google_plan_generation(context): _setup_plan_generation_graph(context) context.response = context.provider.generate_changes( context.project, context.plan, context.contexts, ) context.chat_google_call = context.chat_google_class.call_args context.plan_generation_graph_call = context.plan_generation_graph_class.call_args @then( 'the Google provider should construct ChatGoogleGenerativeAI with api key "{api_key}" and model "{model_id}"' ) def step_assert_google_constructor(context, api_key, model_id): call = getattr(context, "chat_google_call", None) assert call is not None, "ChatGoogleGenerativeAI should have been called" call_args, call_kwargs = call assert call_args == (), "ChatGoogleGenerativeAI should use keyword arguments" assert call_kwargs["api_key"] == api_key assert call_kwargs["model"] == model_id graph_call = getattr(context, "plan_generation_graph_call", None) assert graph_call is not None, "PlanGenerationGraph should receive the LLM" _, graph_kwargs = graph_call assert graph_kwargs["llm"] is context.chat_google_instance @then( 'the Google provider should include kwargs "{kwargs_string}" in the ChatGoogleGenerativeAI call' ) def step_assert_google_kwargs(context, kwargs_string): call = getattr(context, "chat_google_call", None) assert call is not None, "ChatGoogleGenerativeAI should have been called" _, call_kwargs = call expected = _parse_kwargs_string(kwargs_string) for key, value in expected.items(): assert key in call_kwargs, f"Expected {key} in ChatGoogleGenerativeAI kwargs" assert call_kwargs[key] == value, ( f"Expected ChatGoogleGenerativeAI kwargs[{key!r}] to equal {value!r}" ) @then( 'the Google provider metadata should report name "{expected_name}" and model "{expected_model}"' ) def step_assert_google_metadata(context, expected_name, expected_model): provider = getattr(context, "provider", None) assert provider is not None, "Google provider should exist" assert provider.name == expected_name assert provider.model_id == expected_model response = getattr(context, "response", None) if response is not None: assert response.model_used == expected_model @then( 'the Google provider response should include {count:d} generated change for "{file_path}"' ) def step_assert_google_generated_change(context, count, file_path): response = getattr(context, "response", None) assert response is not None, "Provider response should exist" assert len(response.changes) == count, ( f"Expected {count} changes, got {len(response.changes)}" ) assert any( getattr(change, "file_path", None) == file_path for change in response.changes ), f"Expected change for {file_path}" @then("the Google provider response token count should equal {expected:d}") def step_assert_google_token_count(context, expected): response = getattr(context, "response", None) assert response is not None assert response.token_count == expected @then('the Google provider creation should fail with error "{message}"') def step_assert_google_creation_error(context, message): error = getattr(context, "google_provider_error", None) assert error is not None, "Expected the provider to raise an error" assert isinstance(error, ValueError) assert str(error) == message