diff --git a/features/gemini_provider.feature b/features/gemini_provider.feature new file mode 100644 index 000000000..fe8f621c8 --- /dev/null +++ b/features/gemini_provider.feature @@ -0,0 +1,79 @@ +Feature: Gemini provider adapter coverage + As a maintainer integrating real provider adapters + I want unit-level Behave scenarios for the Gemini provider + So that the GeminiProvider implementation stays documented and regression tested + + @unit @providers @gemini + Scenario: Gemini provider instantiates ChatGoogleGenerativeAI with provided credentials + Given I have sample provider domain inputs + When I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + And I request plan generation from the Gemini provider + Then the Gemini provider should construct ChatGoogleGenerativeAI with api key "gemini-unit-key" and model "gemini-1.5-pro" + And the Gemini provider metadata should report name "gemini" and model "gemini-1.5-pro" + + @unit @providers @gemini + Scenario: Gemini provider defaults to gemini-1.5-pro model + Given I have sample provider domain inputs + When I create a Gemini provider with API key "gemini-unit-key" and default model + And I request plan generation from the Gemini provider + Then the Gemini provider metadata should report name "gemini" and model "gemini-1.5-pro" + + @unit @providers @gemini + Scenario: Gemini provider supports gemini-1.5-flash model + Given I have sample provider domain inputs + When I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-flash" + And I request plan generation from the Gemini provider + Then the Gemini provider should construct ChatGoogleGenerativeAI with api key "gemini-unit-key" and model "gemini-1.5-flash" + And the Gemini provider metadata should report name "gemini" and model "gemini-1.5-flash" + + @unit @providers @gemini + Scenario: Gemini provider forwards extra kwargs to ChatGoogleGenerativeAI + Given I have sample provider domain inputs + And I set the Gemini provider extra kwargs "temperature=0.3,max_output_tokens=1024" + When I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + And I request plan generation from the Gemini provider + Then the Gemini provider should construct ChatGoogleGenerativeAI with api key "gemini-unit-key" and model "gemini-1.5-pro" + And the Gemini provider should include kwargs "temperature=0.3,max_output_tokens=1024" in the ChatGoogleGenerativeAI call + + @unit @providers @gemini + Scenario: Gemini provider returns generated changes with token count + Given I have sample provider domain inputs + And the plan generation graph returns a generated change for "app/gemini_provider.py" + And the Gemini provider token estimator returns 512 tokens + When I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + And I request plan generation from the Gemini provider + Then the Gemini provider response should include 1 generated change for "app/gemini_provider.py" + And the Gemini provider response token count should equal 512 + + @unit @providers @gemini + Scenario: Gemini provider rejects missing API key + Given I have sample provider domain inputs + When I attempt to create a Gemini provider without an API key + Then the Gemini provider creation should fail with error "Gemini API key is required" + + @unit @providers @gemini + Scenario: Gemini provider stubbed response reports metadata + Given I have sample provider domain inputs + And I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + When I request plan generation from the Gemini provider + Then the Gemini provider response should contain no generated changes + And the Gemini provider response should report the requested model without errors + + @unit @providers @gemini + Scenario: Gemini provider reports runtime errors + Given I have sample provider domain inputs + And I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + And the plan generation graph raises RuntimeError "gemini quota exceeded" + When I request plan generation from the Gemini provider + Then the Gemini provider response should report error "gemini quota exceeded" + And the Gemini provider response should contain no generated changes + + @unit @providers @gemini @streaming + Scenario: Gemini provider streaming yields workflow events + Given I have sample provider domain inputs + And the plan generation graph returns a generated change for "app/stream_gemini.py" + And the plan generation graph emits streaming nodes "load_context,analyze_requirements,generate_plan,validate" + When I create a Gemini provider with API key "gemini-unit-key" and model "gemini-1.5-pro" + And I stream plan generation from the Gemini provider + Then the Gemini provider streaming events should include nodes "load_context,analyze_requirements,generate_plan,validate" + And the Gemini provider streaming result should finish with a response containing 1 generated change diff --git a/features/steps/gemini_provider_steps.py b/features/steps/gemini_provider_steps.py new file mode 100644 index 000000000..f92afa843 --- /dev/null +++ b/features/steps/gemini_provider_steps.py @@ -0,0 +1,293 @@ +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.gemini_provider import GeminiProvider + + +def _register_cleanup(context: Any, cleanup: Any) -> None: + 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 {} + result: dict[str, Any] = {} + for entry in kwargs_string.split(","): + entry = entry.strip() + if not entry: + continue + if "=" not in entry: + result[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 + result[key] = parsed_value + return result + + +def _setup_plan_generation_graph(context: Any) -> 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: Any, **_kwargs: Any) -> Any: + yield from events_copy + + 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 Gemini provider extra kwargs "{kwargs_string}"') +def step_set_gemini_provider_kwargs(context: Any, kwargs_string: str) -> None: + context.gemini_provider_kwargs = _parse_kwargs_string(kwargs_string) + + +@given("the Gemini provider token estimator returns {token_count:d} tokens") +def step_set_gemini_token_estimator(context: Any, token_count: int) -> None: + context.gemini_token_count = token_count + + +@given( + 'I create a Gemini provider with API key "{api_key}" and model "{model_id}"' +) +@when( + 'I create a Gemini provider with API key "{api_key}" and model "{model_id}"' +) +def step_create_gemini_provider(context: Any, api_key: str, model_id: str) -> None: + patcher = patch("cleveragents.providers.llm.gemini_provider.ChatGoogleGenerativeAI") + context.chat_gemini_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, "gemini_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, "gemini_provider_kwargs", {})) + + context.provider = GeminiProvider( + api_key=api_key, + model=model_id, + **extra_kwargs, + ) + context.chat_gemini_class = mock_chat_class + context.chat_gemini_instance = mock_chat_instance + + +@when('I create a Gemini provider with API key "{api_key}" and default model') +def step_create_gemini_provider_default_model(context: Any, api_key: str) -> None: + patcher = patch("cleveragents.providers.llm.gemini_provider.ChatGoogleGenerativeAI") + context.chat_gemini_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) + mock_chat_class.return_value = mock_chat_instance + + context.provider = GeminiProvider(api_key=api_key) + context.chat_gemini_class = mock_chat_class + context.chat_gemini_instance = mock_chat_instance + + +@when("I attempt to create a Gemini provider without an API key") +def step_gemini_provider_without_api_key(context: Any) -> None: + context.gemini_provider_error = None + try: + GeminiProvider(api_key="", model="gemini-1.5-pro") + except Exception as exc: # pragma: no cover - defensive + context.gemini_provider_error = exc + + +@when("I request plan generation from the Gemini provider") +def step_request_gemini_plan_generation(context: Any) -> None: + _setup_plan_generation_graph(context) + context.response = context.provider.generate_changes( + context.project, + context.plan, + context.contexts, + ) + context.chat_gemini_call = context.chat_gemini_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@when("I stream plan generation from the Gemini provider") +def step_stream_gemini_plan_generation(context: Any) -> None: + _setup_plan_generation_graph(context) + context.streamed_events = list( + context.provider.stream_changes( + context.project, + context.plan, + context.contexts, + ) + ) + context.chat_gemini_call = context.chat_gemini_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@then( + 'the Gemini provider should construct ChatGoogleGenerativeAI with api key "{api_key}" and model "{model_id}"' +) +def step_assert_gemini_constructor( + context: Any, api_key: str, model_id: str +) -> None: + call = getattr(context, "chat_gemini_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_gemini_instance + + +@then( + 'the Gemini provider should include kwargs "{kwargs_string}" in the ChatGoogleGenerativeAI call' +) +def step_assert_gemini_kwargs(context: Any, kwargs_string: str) -> None: + call = getattr(context, "chat_gemini_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 Gemini provider metadata should report name "{expected_name}" and model "{expected_model}"' +) +def step_assert_gemini_metadata( + context: Any, expected_name: str, expected_model: str +) -> None: + provider = getattr(context, "provider", None) + assert provider is not None, "Gemini 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 Gemini provider response should include {count:d} generated change for "{file_path}"' +) +def step_assert_gemini_generated_change( + context: Any, count: int, file_path: str +) -> None: + 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 Gemini provider response token count should equal {expected:d}") +def step_assert_gemini_token_count(context: Any, expected: int) -> None: + response = getattr(context, "response", None) + assert response is not None + assert response.token_count == expected + + +@then('the Gemini provider creation should fail with error "{message}"') +def step_assert_gemini_creation_error(context: Any, message: str) -> None: + error = getattr(context, "gemini_provider_error", None) + assert error is not None, "Expected the provider to raise an error" + assert isinstance(error, ValueError) + assert str(error) == message + + +@then("the Gemini provider response should contain no generated changes") +def step_assert_gemini_no_changes(context: Any) -> None: + response = getattr(context, "response", None) + assert response is not None, "Provider response should exist" + assert len(response.changes) == 0, "Expected no generated changes" + + +@then("the Gemini provider response should report the requested model without errors") +def step_assert_gemini_placeholder_metadata(context: Any) -> None: + assert context.response is not None + assert context.response.model_used == context.provider.model_id + assert context.response.error_message in (None, "") + + +@then('the Gemini provider response should report error "{message}"') +def step_assert_gemini_response_error(context: Any, message: str) -> None: + assert context.response is not None + assert context.response.error_message == message + + +@then('the Gemini provider streaming events should include nodes "{node_list}"') +def step_assert_gemini_stream_events(context: Any, node_list: str) -> None: + expected = [node.strip() for node in node_list.split(",") if node.strip()] + actual = [ + next(iter(event.keys())) + for event in getattr(context, "streamed_events", []) + if "__end__" not in event + ] + assert actual == expected + + +@then( + "the Gemini provider streaming result should finish with a response containing {count:d} generated change" +) +def step_assert_gemini_stream_final_response(context: Any, count: int) -> None: + events = getattr(context, "streamed_events", []) + assert events, "Expected streamed events" + final_event = events[-1] + assert "__end__" in final_event, "Expected final __end__ event" + response = final_event["__end__"].get("response") + assert response is not None, "Expected ProviderResponse in __end__ event" + assert len(response.changes) == count diff --git a/src/cleveragents/providers/llm/__init__.py b/src/cleveragents/providers/llm/__init__.py index 9a07b2bed..af4e055f4 100644 --- a/src/cleveragents/providers/llm/__init__.py +++ b/src/cleveragents/providers/llm/__init__.py @@ -5,12 +5,14 @@ import them directly from the ``cleveragents.providers.llm`` subpackage. """ from cleveragents.providers.llm.anthropic_provider import AnthropicChatProvider +from cleveragents.providers.llm.gemini_provider import GeminiProvider from cleveragents.providers.llm.google_provider import GoogleChatProvider from cleveragents.providers.llm.openai_provider import OpenAIChatProvider from cleveragents.providers.llm.openrouter_provider import OpenRouterChatProvider __all__ = [ "AnthropicChatProvider", + "GeminiProvider", "GoogleChatProvider", "OpenAIChatProvider", "OpenRouterChatProvider", diff --git a/src/cleveragents/providers/llm/gemini_provider.py b/src/cleveragents/providers/llm/gemini_provider.py new file mode 100644 index 000000000..dd77c04c3 --- /dev/null +++ b/src/cleveragents/providers/llm/gemini_provider.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +from typing import Any + +from langchain_google_genai import ChatGoogleGenerativeAI + +from cleveragents.providers.llm.langchain_chat_provider import LangChainChatProvider + + +class GeminiProvider(LangChainChatProvider): + """Concrete provider adapter for Google Gemini models via the Gemini API. + + Uses the ``GEMINI_API_KEY`` environment variable (or ``gemini_api_key`` + settings field) rather than the generic ``GOOGLE_API_KEY`` used by + :class:`~cleveragents.providers.llm.google_provider.GoogleChatProvider`. + + Supports Gemini 1.5 Pro, Gemini 1.5 Flash, and other Gemini model + variants with full streaming and tool-calling capabilities. + + Example:: + + import os + provider = GeminiProvider(api_key=os.environ["GEMINI_API_KEY"]) + response = provider.generate_changes(project, plan, contexts) + """ + + def __init__( + self, + *, + api_key: str, + model: str = "gemini-1.5-pro", + max_retries: int = 3, + **llm_kwargs: Any, + ) -> None: + """Initialise the Gemini provider. + + Args: + api_key: Google Gemini API key. Must be non-empty. + model: Gemini model identifier. Defaults to ``gemini-1.5-pro``. + Other supported values include ``gemini-1.5-flash``. + max_retries: Maximum number of retry attempts on transient errors. + **llm_kwargs: Additional keyword arguments forwarded verbatim to + :class:`~langchain_google_genai.ChatGoogleGenerativeAI`. + + Raises: + ValueError: When *api_key* is empty or ``None``. + """ + if not api_key: + raise ValueError("Gemini API key is required") + + def factory(resolved_model: str) -> ChatGoogleGenerativeAI: + kwargs: dict[str, Any] = { + "api_key": api_key, + "model": resolved_model, + } + if llm_kwargs: + kwargs.update(llm_kwargs) + return ChatGoogleGenerativeAI(**kwargs) + + super().__init__( + name="gemini", + model_id=model, + llm_factory=factory, + max_retries=max_retries, + supports_streaming=True, + ) diff --git a/src/cleveragents/providers/registry.py b/src/cleveragents/providers/registry.py index 45943375c..6b9454399 100644 --- a/src/cleveragents/providers/registry.py +++ b/src/cleveragents/providers/registry.py @@ -664,6 +664,27 @@ class ProviderRegistry: max_retries=max_retries, ) + if provider_type == ProviderType.GEMINI: + from cleveragents.providers.llm.gemini_provider import GeminiProvider + + key_attr = self.PROVIDER_KEY_ATTRS.get(provider_type) + api_key = getattr(self._settings, key_attr, None) if key_attr else None + if not api_key: + missing_env = ( + key_attr.upper() if key_attr else provider_type.value.upper() + ) + raise ValueError( + f"Provider {provider_type.value} is not configured. " + f"Please set the {missing_env} environment variable." + ) + + return GeminiProvider( + api_key=api_key, + model=model_id + or self.DEFAULT_MODELS.get(provider_type, "gemini-1.5-pro"), + max_retries=max_retries, + ) + if provider_type == ProviderType.OPENROUTER: from cleveragents.providers.llm.openrouter_provider import ( OpenRouterChatProvider,