feat(providers): implement GeminiProvider for Google Gemini API #10617
@@ -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
|
||||
@@ -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``.
|
||||
|
HAL9001
commented
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,
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user
TEST: Kwargs scenario only checks _llm_kwargs, not actual generation config params.