From cb1ceda79cb2d5ad108ef2fd6f6e704ac61b13a2 Mon Sep 17 00:00:00 2001 From: HAL9000 Date: Sat, 18 Apr 2026 19:18:36 +0000 Subject: [PATCH] feat(providers): implement OllamaProvider and MistralProvider - Implemented OllamaChatProvider to enable local Ollama model support. - Implemented MistralChatProvider to integrate with the Mistral API. - Added Behave BDD tests for both providers. - Updated dependencies: langchain-mistralai and ollama. - Updated provider exports to include the new providers. ISSUES CLOSED: #5257 --- features/mistral_provider.feature | 71 ++++ features/ollama_provider.feature | 70 ++++ features/steps/mistral_provider_steps.py | 333 ++++++++++++++++++ features/steps/ollama_provider_steps.py | 321 +++++++++++++++++ pyproject.toml | 2 + src/cleveragents/providers/llm/__init__.py | 4 + .../providers/llm/mistral_provider.py | 46 +++ .../providers/llm/ollama_provider.py | 37 ++ 8 files changed, 884 insertions(+) create mode 100644 features/mistral_provider.feature create mode 100644 features/ollama_provider.feature create mode 100644 features/steps/mistral_provider_steps.py create mode 100644 features/steps/ollama_provider_steps.py create mode 100644 src/cleveragents/providers/llm/mistral_provider.py create mode 100644 src/cleveragents/providers/llm/ollama_provider.py diff --git a/features/mistral_provider.feature b/features/mistral_provider.feature new file mode 100644 index 000000000..ef788522d --- /dev/null +++ b/features/mistral_provider.feature @@ -0,0 +1,71 @@ +Feature: Mistral chat provider coverage + As a maintainer integrating Mistral API support + I want unit-level Behave scenarios for the Mistral chat provider + So that the Mistral adapter stays documented and regression tested + + @unit @providers @mistral + Scenario: Mistral provider instantiates ChatMistralAI with provided credentials + Given I have sample provider domain inputs + When I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + And I request plan generation from the Mistral provider + Then the Mistral provider should construct ChatMistralAI with api key "test-key-123" and model "mistral-large-latest" + And the Mistral provider metadata should report name "mistral" and model "mistral-large-latest" + + @unit @providers @mistral + Scenario: Mistral provider stubbed response reports metadata + Given I have sample provider domain inputs + And I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + When I request plan generation from the Mistral provider + Then the Mistral provider response should contain no generated changes + And the Mistral provider response should report the requested model without errors + + @unit @providers @mistral + Scenario: Mistral provider rejects missing API key + Given I have sample provider domain inputs + When I attempt to create a Mistral chat provider without an API key + Then the Mistral provider creation should fail with error containing "Mistral API key is required" + + @unit @providers @mistral + Scenario: Mistral provider reads API key from environment variable + Given I have sample provider domain inputs + And I set the MISTRAL_API_KEY environment variable to "env-key-456" + When I create a Mistral chat provider without explicit API key and model "mistral-large-latest" + And I request plan generation from the Mistral provider + Then the Mistral provider should construct ChatMistralAI with api key "env-key-456" and model "mistral-large-latest" + + @unit @providers @mistral + Scenario: Mistral provider forwards extra kwargs + Given I have sample provider domain inputs + And I set the Mistral provider extra kwargs "temperature=0.5,max_tokens=512" + When I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + And I request plan generation from the Mistral provider + Then the Mistral provider should construct ChatMistralAI with api key "test-key-123" and model "mistral-large-latest" + And the Mistral provider should include kwargs "temperature=0.5,max_tokens=512" in the ChatMistralAI call + + @unit @providers @mistral + Scenario: Mistral provider reports runtime errors + Given I have sample provider domain inputs + And I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + And the plan generation graph raises RuntimeError "API rate limit exceeded" + When I request plan generation from the Mistral provider + Then the Mistral provider response should report error "API rate limit exceeded" + And the Mistral provider response should contain no generated changes + + @unit @providers @mistral + Scenario: Mistral provider streaming yields workflow events + Given I have sample provider domain inputs + And the plan generation graph returns a generated change for "app/mistral.py" + And the plan generation graph emits streaming nodes "load_context,analyze_requirements,generate_plan,validate" + And I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + When I stream plan generation from the Mistral provider + Then the Mistral provider streaming events should include nodes "load_context,analyze_requirements,generate_plan,validate" + And the Mistral provider streaming result should finish with a response containing 1 generated change + + @unit @providers @mistral + Scenario: Mistral provider surfaces plan generation errors + Given I have sample provider domain inputs + And the plan generation graph raises ValueError "invalid request" + And I create a Mistral chat provider with API key "test-key-123" and model "mistral-large-latest" + When I request plan generation from the Mistral provider + Then the Mistral provider response should report error "invalid request" + And the Mistral provider response should contain no generated changes diff --git a/features/ollama_provider.feature b/features/ollama_provider.feature new file mode 100644 index 000000000..386f153fb --- /dev/null +++ b/features/ollama_provider.feature @@ -0,0 +1,70 @@ +Feature: Ollama chat provider coverage + As a maintainer integrating local model support + I want unit-level Behave scenarios for the Ollama chat provider + So that the local model adapter stays documented and regression tested + + @unit @providers @ollama + Scenario: Ollama provider instantiates ChatOllama with provided credentials + Given I have sample provider domain inputs + When I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + And I request plan generation from the Ollama provider + Then the Ollama provider should construct ChatOllama with model "llama2" and base_url "http://localhost:11434" + And the Ollama provider metadata should report name "ollama" and model "llama2" + + @unit @providers @ollama + Scenario: Ollama provider stubbed response reports metadata + Given I have sample provider domain inputs + And I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + When I request plan generation from the Ollama provider + Then the Ollama provider response should contain no generated changes + And the Ollama provider response should report the requested model without errors + + @unit @providers @ollama + Scenario: Ollama provider rejects missing model name + Given I have sample provider domain inputs + When I attempt to create an Ollama chat provider without a model name + Then the Ollama provider creation should fail with error "Ollama model name is required" + + @unit @providers @ollama + Scenario: Ollama provider uses default base URL + Given I have sample provider domain inputs + When I create an Ollama chat provider with model "llama2" and default base_url + And I request plan generation from the Ollama provider + Then the Ollama provider should construct ChatOllama with model "llama2" and base_url "http://localhost:11434" + + @unit @providers @ollama + Scenario: Ollama provider forwards extra kwargs + Given I have sample provider domain inputs + And I set the Ollama provider extra kwargs "temperature=0.7,top_p=0.9" + When I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + And I request plan generation from the Ollama provider + Then the Ollama provider should construct ChatOllama with model "llama2" and base_url "http://localhost:11434" + And the Ollama provider should include kwargs "temperature=0.7,top_p=0.9" in the ChatOllama call + + @unit @providers @ollama + Scenario: Ollama provider reports runtime errors + Given I have sample provider domain inputs + And I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + And the plan generation graph raises RuntimeError "connection refused" + When I request plan generation from the Ollama provider + Then the Ollama provider response should report error "connection refused" + And the Ollama provider response should contain no generated changes + + @unit @providers @ollama + Scenario: Ollama provider streaming yields workflow events + Given I have sample provider domain inputs + And the plan generation graph returns a generated change for "app/stream.py" + And the plan generation graph emits streaming nodes "load_context,analyze_requirements,generate_plan,validate" + And I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + When I stream plan generation from the Ollama provider + Then the Ollama provider streaming events should include nodes "load_context,analyze_requirements,generate_plan,validate" + And the Ollama provider streaming result should finish with a response containing 1 generated change + + @unit @providers @ollama + Scenario: Ollama provider surfaces plan generation errors + Given I have sample provider domain inputs + And the plan generation graph raises ValueError "model not found" + And I create an Ollama chat provider with model "llama2" and base_url "http://localhost:11434" + When I request plan generation from the Ollama provider + Then the Ollama provider response should report error "model not found" + And the Ollama provider response should contain no generated changes diff --git a/features/steps/mistral_provider_steps.py b/features/steps/mistral_provider_steps.py new file mode 100644 index 000000000..38707435c --- /dev/null +++ b/features/steps/mistral_provider_steps.py @@ -0,0 +1,333 @@ +from __future__ import annotations + +import ast +import os +from typing import Any +from unittest.mock import MagicMock, patch + +from behave import given, then, when + +from cleveragents.domain.models.core import ( + Context, + OperationType, + Plan, + Project, +) +from cleveragents.providers.llm.mistral_provider import MistralChatProvider + + +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 {} + 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) -> 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): + 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 have sample provider domain inputs") +def step_sample_provider_inputs(context): + context.project = MagicMock(spec=Project) + context.plan = MagicMock(spec=Plan) + context.plan.prompt = "Add placeholder coverage" + context.contexts = [MagicMock(spec=Context)] + context.contexts[0].content = "Sample context entry" + + +@given('the plan generation graph returns a generated change for "{file_path}"') +def step_plan_generation_returns_change(context, file_path): + change = { + "plan_id": 1, + "file_path": file_path, + "operation": OperationType.MODIFY.value, + "new_content": "# updated content", + } + context.plan_generation_state_override = { + "generated_changes": [change], + "validation_result": {"status": "PASS"}, + "error": None, + } + + +@given('the plan generation graph emits streaming nodes "{node_list}"') +def step_plan_generation_stream_nodes(context, node_list): + nodes = [node.strip() for node in node_list.split(",") if node.strip()] + state = getattr(context, "plan_generation_state_override", {}) or {} + events: list[dict[str, Any]] = [] + for node in nodes: + payload: dict[str, Any] = {"status": "completed"} + if node == "generate_plan" and isinstance(state.get("generated_changes"), list): + payload = { + "generated_changes": state["generated_changes"], + "status": "completed", + } + elif node == "validate" and isinstance(state.get("validation_result"), dict): + payload = { + "validation_result": state["validation_result"], + "status": "completed", + } + events.append({node: payload}) + context.plan_generation_stream_events = events + + +@given('the plan generation graph raises ValueError "{message}"') +def step_plan_generation_raises_value_error(context, message): + context.plan_generation_invoke_side_effect = ValueError(message) + + +@given('the plan generation graph raises RuntimeError "{message}"') +def step_plan_generation_runtime_error(context, message): + context.plan_generation_invoke_side_effect = RuntimeError(message) + + +@given('I set the MISTRAL_API_KEY environment variable to "{api_key}"') +def step_set_mistral_env_var(context, api_key): + context.mistral_env_api_key = api_key + os.environ["MISTRAL_API_KEY"] = api_key + _register_cleanup(context, lambda: os.environ.pop("MISTRAL_API_KEY", None)) + + +@given('I set the Mistral provider extra kwargs "{kwargs_string}"') +def step_set_mistral_provider_kwargs(context, kwargs_string): + context.mistral_provider_kwargs = _parse_kwargs_string(kwargs_string) + + +@given( + 'I create a Mistral chat provider with API key "{api_key}" and model "{model}"' +) +@when( + 'I create a Mistral chat provider with API key "{api_key}" and model "{model}"' +) +def step_create_mistral_provider(context, api_key, model): + patcher = patch("cleveragents.providers.llm.mistral_provider.ChatMistralAI") + context.chat_mistral_patcher = patcher + mock_chat_mistral_class = patcher.start() + _register_cleanup(context, patcher.stop) + mock_chat_mistral_instance = MagicMock(name="ChatMistralAIInstance") + mock_chat_mistral_instance.get_num_tokens = MagicMock(return_value=0) + mock_chat_mistral_class.return_value = mock_chat_mistral_instance + + extra_kwargs = dict(getattr(context, "mistral_provider_kwargs", {})) + + context.provider = MistralChatProvider( + api_key=api_key, + model=model, + **extra_kwargs, + ) + context.chat_mistral_class = mock_chat_mistral_class + context.chat_mistral_instance = mock_chat_mistral_instance + + +@when( + 'I create a Mistral chat provider without explicit API key and model "{model}"' +) +def step_create_mistral_provider_from_env(context, model): + patcher = patch("cleveragents.providers.llm.mistral_provider.ChatMistralAI") + context.chat_mistral_patcher = patcher + mock_chat_mistral_class = patcher.start() + _register_cleanup(context, patcher.stop) + mock_chat_mistral_instance = MagicMock(name="ChatMistralAIInstance") + mock_chat_mistral_instance.get_num_tokens = MagicMock(return_value=0) + mock_chat_mistral_class.return_value = mock_chat_mistral_instance + + extra_kwargs = dict(getattr(context, "mistral_provider_kwargs", {})) + + context.provider = MistralChatProvider( + model=model, + **extra_kwargs, + ) + context.chat_mistral_class = mock_chat_mistral_class + context.chat_mistral_instance = mock_chat_mistral_instance + + +@when("I attempt to create a Mistral chat provider without an API key") +def step_mistral_provider_without_api_key(context): + context.mistral_provider_error = None + # Make sure env var is not set + os.environ.pop("MISTRAL_API_KEY", None) + try: + MistralChatProvider(api_key="", model="mistral-large-latest") + except Exception as exc: # pragma: no cover - defensive logging only + context.mistral_provider_error = exc + + +@when("I request plan generation from the Mistral provider") +def step_request_plan_generation(context): + _setup_plan_generation_graph(context) + + context.response = context.provider.generate_changes( + context.project, + context.plan, + context.contexts, + ) + context.chat_mistral_call = context.chat_mistral_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@when("I stream plan generation from the Mistral provider") +def step_stream_plan_generation(context): + _setup_plan_generation_graph(context) + context.streamed_events = list( + context.provider.stream_changes( + context.project, + context.plan, + context.contexts, + ) + ) + context.chat_mistral_call = context.chat_mistral_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@then( + 'the Mistral provider should construct ChatMistralAI with api key "{api_key}" and model "{model}"' +) +def step_assert_chat_mistral_constructor(context, api_key, model): + assert context.chat_mistral_call is not None, "ChatMistralAI should have been called" + call_args, call_kwargs = context.chat_mistral_call + assert call_args == (), "ChatMistralAI should be called with keyword arguments" + assert call_kwargs["api_key"] == api_key + assert call_kwargs["model"] == model + + assert context.plan_generation_graph_call is not None, ( + "PlanGenerationGraph should receive the ChatMistralAI instance" + ) + _, graph_kwargs = context.plan_generation_graph_call + assert graph_kwargs["llm"] is context.chat_mistral_instance + + +@then( + 'the Mistral provider should include kwargs "{kwargs_string}" in the ChatMistralAI call' +) +def step_assert_chat_mistral_kwargs(context, kwargs_string): + assert context.chat_mistral_call is not None, "ChatMistralAI should have been called" + _, call_kwargs = context.chat_mistral_call + expected_kwargs = _parse_kwargs_string(kwargs_string) + for key, value in expected_kwargs.items(): + assert key in call_kwargs, f"Expected {key} in ChatMistralAI kwargs" + assert call_kwargs[key] == value, ( + f"Expected ChatMistralAI kwargs[{key!r}] to equal {value!r}" + ) + + +@then( + 'the Mistral provider metadata should report name "{expected_name}" and model "{expected_model}"' +) +def step_assert_provider_metadata(context, expected_name, expected_model): + provider = getattr(context, "provider", None) + assert provider is not None, "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 Mistral provider response should report the requested model without errors") +def step_assert_placeholder_metadata(context): + assert context.response is not None + assert context.response.model_used == context.provider.model_id + assert context.response.error_message in (None, "") + + +@then("the Mistral provider response should contain no generated changes") +def step_assert_no_changes(context): + assert context.response is not None, "Provider response should exist" + assert len(context.response.changes) == 0, "Expected no generated changes" + + +@then('the Mistral provider response should report error "{message}"') +def step_assert_response_error(context, message): + assert context.response is not None + assert context.response.error_message == message + + +@then('the Mistral provider streaming events should include nodes "{node_list}"') +def step_assert_stream_events(context, node_list): + 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 Mistral provider streaming result should finish with a response containing {count:d} generated change" +) +def step_assert_stream_final_response(context, count): + 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 + + +@then('the Mistral provider creation should fail with error containing "{message}"') +def step_assert_mistral_provider_error(context, message): + error = getattr(context, "mistral_provider_error", None) + assert error is not None, "Expected the provider to raise an error" + assert isinstance(error, ValueError) + assert message in str(error) diff --git a/features/steps/ollama_provider_steps.py b/features/steps/ollama_provider_steps.py new file mode 100644 index 000000000..792439d88 --- /dev/null +++ b/features/steps/ollama_provider_steps.py @@ -0,0 +1,321 @@ +from __future__ import annotations + +import ast +from typing import Any +from unittest.mock import MagicMock, patch + +from behave import given, then, when + +from cleveragents.domain.models.core import ( + Context, + OperationType, + Plan, + Project, +) +from cleveragents.providers.llm.ollama_provider import OllamaChatProvider + + +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 {} + 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) -> 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): + 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 have sample provider domain inputs") +def step_sample_provider_inputs(context): + context.project = MagicMock(spec=Project) + context.plan = MagicMock(spec=Plan) + context.plan.prompt = "Add placeholder coverage" + context.contexts = [MagicMock(spec=Context)] + context.contexts[0].content = "Sample context entry" + + +@given('the plan generation graph returns a generated change for "{file_path}"') +def step_plan_generation_returns_change(context, file_path): + change = { + "plan_id": 1, + "file_path": file_path, + "operation": OperationType.MODIFY.value, + "new_content": "# updated content", + } + context.plan_generation_state_override = { + "generated_changes": [change], + "validation_result": {"status": "PASS"}, + "error": None, + } + + +@given('the plan generation graph emits streaming nodes "{node_list}"') +def step_plan_generation_stream_nodes(context, node_list): + nodes = [node.strip() for node in node_list.split(",") if node.strip()] + state = getattr(context, "plan_generation_state_override", {}) or {} + events: list[dict[str, Any]] = [] + for node in nodes: + payload: dict[str, Any] = {"status": "completed"} + if node == "generate_plan" and isinstance(state.get("generated_changes"), list): + payload = { + "generated_changes": state["generated_changes"], + "status": "completed", + } + elif node == "validate" and isinstance(state.get("validation_result"), dict): + payload = { + "validation_result": state["validation_result"], + "status": "completed", + } + events.append({node: payload}) + context.plan_generation_stream_events = events + + +@given('the plan generation graph raises ValueError "{message}"') +def step_plan_generation_raises_value_error(context, message): + context.plan_generation_invoke_side_effect = ValueError(message) + + +@given('the plan generation graph raises RuntimeError "{message}"') +def step_plan_generation_runtime_error(context, message): + context.plan_generation_invoke_side_effect = RuntimeError(message) + + +@given('I set the Ollama provider extra kwargs "{kwargs_string}"') +def step_set_ollama_provider_kwargs(context, kwargs_string): + context.ollama_provider_kwargs = _parse_kwargs_string(kwargs_string) + + +@given( + 'I create an Ollama chat provider with model "{model}" and base_url "{base_url}"' +) +@when( + 'I create an Ollama chat provider with model "{model}" and base_url "{base_url}"' +) +def step_create_ollama_provider(context, model, base_url): + patcher = patch("cleveragents.providers.llm.ollama_provider.ChatOllama") + context.chat_ollama_patcher = patcher + mock_chat_ollama_class = patcher.start() + _register_cleanup(context, patcher.stop) + mock_chat_ollama_instance = MagicMock(name="ChatOllamaInstance") + mock_chat_ollama_instance.get_num_tokens = MagicMock(return_value=0) + mock_chat_ollama_class.return_value = mock_chat_ollama_instance + + extra_kwargs = dict(getattr(context, "ollama_provider_kwargs", {})) + + context.provider = OllamaChatProvider( + model=model, + base_url=base_url, + **extra_kwargs, + ) + context.chat_ollama_class = mock_chat_ollama_class + context.chat_ollama_instance = mock_chat_ollama_instance + + +@when('I create an Ollama chat provider with model "{model}" and default base_url') +def step_create_ollama_provider_default_url(context, model): + patcher = patch("cleveragents.providers.llm.ollama_provider.ChatOllama") + context.chat_ollama_patcher = patcher + mock_chat_ollama_class = patcher.start() + _register_cleanup(context, patcher.stop) + mock_chat_ollama_instance = MagicMock(name="ChatOllamaInstance") + mock_chat_ollama_instance.get_num_tokens = MagicMock(return_value=0) + mock_chat_ollama_class.return_value = mock_chat_ollama_instance + + extra_kwargs = dict(getattr(context, "ollama_provider_kwargs", {})) + + context.provider = OllamaChatProvider( + model=model, + **extra_kwargs, + ) + context.chat_ollama_class = mock_chat_ollama_class + context.chat_ollama_instance = mock_chat_ollama_instance + + +@when("I attempt to create an Ollama chat provider without a model name") +def step_ollama_provider_without_model(context): + context.ollama_provider_error = None + try: + OllamaChatProvider(model="") + except Exception as exc: # pragma: no cover - defensive logging only + context.ollama_provider_error = exc + + +@when("I request plan generation from the Ollama provider") +def step_request_plan_generation(context): + _setup_plan_generation_graph(context) + + context.response = context.provider.generate_changes( + context.project, + context.plan, + context.contexts, + ) + context.chat_ollama_call = context.chat_ollama_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@when("I stream plan generation from the Ollama provider") +def step_stream_plan_generation(context): + _setup_plan_generation_graph(context) + context.streamed_events = list( + context.provider.stream_changes( + context.project, + context.plan, + context.contexts, + ) + ) + context.chat_ollama_call = context.chat_ollama_class.call_args + context.plan_generation_graph_call = context.plan_generation_graph_class.call_args + + +@then( + 'the Ollama provider should construct ChatOllama with model "{model}" and base_url "{base_url}"' +) +def step_assert_chat_ollama_constructor(context, model, base_url): + assert context.chat_ollama_call is not None, "ChatOllama should have been called" + call_args, call_kwargs = context.chat_ollama_call + assert call_args == (), "ChatOllama should be called with keyword arguments" + assert call_kwargs["model"] == model + assert call_kwargs["base_url"] == base_url + + assert context.plan_generation_graph_call is not None, ( + "PlanGenerationGraph should receive the ChatOllama instance" + ) + _, graph_kwargs = context.plan_generation_graph_call + assert graph_kwargs["llm"] is context.chat_ollama_instance + + +@then( + 'the Ollama provider should include kwargs "{kwargs_string}" in the ChatOllama call' +) +def step_assert_chat_ollama_kwargs(context, kwargs_string): + assert context.chat_ollama_call is not None, "ChatOllama should have been called" + _, call_kwargs = context.chat_ollama_call + expected_kwargs = _parse_kwargs_string(kwargs_string) + for key, value in expected_kwargs.items(): + assert key in call_kwargs, f"Expected {key} in ChatOllama kwargs" + assert call_kwargs[key] == value, ( + f"Expected ChatOllama kwargs[{key!r}] to equal {value!r}" + ) + + +@then( + 'the Ollama provider metadata should report name "{expected_name}" and model "{expected_model}"' +) +def step_assert_provider_metadata(context, expected_name, expected_model): + provider = getattr(context, "provider", None) + assert provider is not None, "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 Ollama provider response should report the requested model without errors") +def step_assert_placeholder_metadata(context): + assert context.response is not None + assert context.response.model_used == context.provider.model_id + assert context.response.error_message in (None, "") + + +@then("the Ollama provider response should contain no generated changes") +def step_assert_no_changes(context): + assert context.response is not None, "Provider response should exist" + assert len(context.response.changes) == 0, "Expected no generated changes" + + +@then('the Ollama provider response should report error "{message}"') +def step_assert_response_error(context, message): + assert context.response is not None + assert context.response.error_message == message + + +@then('the Ollama provider streaming events should include nodes "{node_list}"') +def step_assert_stream_events(context, node_list): + 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 Ollama provider streaming result should finish with a response containing {count:d} generated change" +) +def step_assert_stream_final_response(context, count): + 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 + + +@then('the Ollama provider creation should fail with error "{message}"') +def step_assert_ollama_provider_error(context, message): + error = getattr(context, "ollama_provider_error", None) + assert error is not None, "Expected the provider to raise an error" + assert isinstance(error, ValueError) + assert str(error) == message diff --git a/pyproject.toml b/pyproject.toml index 09960045b..3c30255ee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,6 +39,7 @@ dependencies = [ "langchain-community>=0.2.14", "langchain-openai>=0.2.0", "langchain-google-genai>=0.2.0", + "langchain-mistralai>=0.1.0", # Mistral API integration "jinja2>=3.1.0", "alembic>=1.13.1", "numpy>=2.1.0", @@ -50,6 +51,7 @@ dependencies = [ "aiohttp>=3.13.4", # CVE-2026-34515 mitigation: open redirect vulnerability "pyyaml>=6.0.3", # Security: address known YAML parsing vulnerabilities "a2a-sdk>=0.3.0,<1.0.0", # A2A Python SDK — required transport for local (stdio) and server (HTTP) modes (ADR-047); pinned <1.0.0 (removed legacy A2AClient) + "ollama>=0.1.0", # Ollama local model support ] [project.optional-dependencies] diff --git a/src/cleveragents/providers/llm/__init__.py b/src/cleveragents/providers/llm/__init__.py index 9a07b2bed..489744d87 100644 --- a/src/cleveragents/providers/llm/__init__.py +++ b/src/cleveragents/providers/llm/__init__.py @@ -6,12 +6,16 @@ import them directly from the ``cleveragents.providers.llm`` subpackage. from cleveragents.providers.llm.anthropic_provider import AnthropicChatProvider from cleveragents.providers.llm.google_provider import GoogleChatProvider +from cleveragents.providers.llm.mistral_provider import MistralChatProvider +from cleveragents.providers.llm.ollama_provider import OllamaChatProvider from cleveragents.providers.llm.openai_provider import OpenAIChatProvider from cleveragents.providers.llm.openrouter_provider import OpenRouterChatProvider __all__ = [ "AnthropicChatProvider", "GoogleChatProvider", + "MistralChatProvider", + "OllamaChatProvider", "OpenAIChatProvider", "OpenRouterChatProvider", ] diff --git a/src/cleveragents/providers/llm/mistral_provider.py b/src/cleveragents/providers/llm/mistral_provider.py new file mode 100644 index 000000000..5fd022c90 --- /dev/null +++ b/src/cleveragents/providers/llm/mistral_provider.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +import os +from typing import Any, cast + +from langchain_core.language_models import BaseLanguageModel +from langchain_mistralai import ChatMistralAI + +from cleveragents.providers.llm.langchain_chat_provider import LangChainChatProvider + + +class MistralChatProvider(LangChainChatProvider): + """Concrete LangChain-backed provider for Mistral API models.""" + + def __init__( + self, + *, + api_key: str | None = None, + model: str = "mistral-large-latest", + max_retries: int = 3, + **llm_kwargs: Any, + ) -> None: + # Use provided API key or fall back to environment variable + resolved_api_key = api_key or os.environ.get("MISTRAL_API_KEY") + if not resolved_api_key: + raise ValueError( + "Mistral API key is required. " + "Provide it via api_key parameter or MISTRAL_API_KEY env var" + ) + + def factory(resolved_model: str) -> BaseLanguageModel: + kwargs: dict[str, Any] = { + "api_key": resolved_api_key, + "model": resolved_model, + } + if llm_kwargs: + kwargs.update(llm_kwargs) + return cast(BaseLanguageModel, ChatMistralAI(**kwargs)) + + super().__init__( + name="mistral", + model_id=model, + llm_factory=factory, + max_retries=max_retries, + supports_streaming=True, + ) diff --git a/src/cleveragents/providers/llm/ollama_provider.py b/src/cleveragents/providers/llm/ollama_provider.py new file mode 100644 index 000000000..4f313215e --- /dev/null +++ b/src/cleveragents/providers/llm/ollama_provider.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +from typing import Any, cast + +from langchain_core.language_models import BaseLanguageModel +from langchain_community.chat_models import ChatOllama + +from cleveragents.providers.llm.langchain_chat_provider import LangChainChatProvider + + +class OllamaChatProvider(LangChainChatProvider): + """Concrete LangChain-backed provider for Ollama local models.""" + + def __init__( + self, + *, + model: str = "llama2", + base_url: str = "http://localhost:11434", + max_retries: int = 3, + **llm_kwargs: Any, + ) -> None: + if not model: + raise ValueError("Ollama model name is required") + + def factory(resolved_model: str) -> BaseLanguageModel: + kwargs: dict[str, Any] = {"model": resolved_model, "base_url": base_url} + if llm_kwargs: + kwargs.update(llm_kwargs) + return cast(BaseLanguageModel, ChatOllama(**kwargs)) + + super().__init__( + name="ollama", + model_id=model, + llm_factory=factory, + max_retries=max_retries, + supports_streaming=True, + )