feat(providers): implement GeminiProvider for Google Gemini API #10617

Closed
HAL9000 wants to merge 1 commits from feat/v3.6.0/gemini-provider into master
5 changed files with 461 additions and 0 deletions
+79
View File
@@ -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
Outdated
Review

TEST: Kwargs scenario only checks _llm_kwargs, not actual generation config params.

TEST: Kwargs scenario only checks _llm_kwargs, not actual generation config params.
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
+293
View File
@@ -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
@@ -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",
@@ -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``.
Outdated
Review

SUGGESTION: _ensure_genai_imported guard only checks GenerativeModel. Add comment that modules load as a group.

SUGGESTION: _ensure_genai_imported guard only checks GenerativeModel. Add comment that modules load as a group.
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,
)
+21
View File
@@ -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,