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.openai_provider import OpenAIChatProvider 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 OpenAI provider token estimator returns {token_count:d} tokens") def step_openai_token_estimator(context, token_count): context.openai_token_count = token_count @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 OpenAI provider organization "{organization}"') def step_set_openai_provider_org(context, organization): if organization.lower() == "none": context.openai_provider_org = None else: context.openai_provider_org = organization @given('I set the OpenAI provider extra kwargs "{kwargs_string}"') def step_set_openai_provider_kwargs(context, kwargs_string): context.openai_provider_kwargs = _parse_kwargs_string(kwargs_string) @given( 'I create an OpenAI chat provider with API key "{api_key}" and model "{model_id}"' ) @when( 'I create an OpenAI chat provider with API key "{api_key}" and model "{model_id}"' ) def step_create_openai_provider(context, api_key, model_id): patcher = patch("cleveragents.providers.llm.openai_provider.ChatOpenAI") context.chat_openai_patcher = patcher mock_chat_openai_class = patcher.start() _register_cleanup(context, patcher.stop) mock_chat_openai_instance = MagicMock(name="ChatOpenAIInstance") mock_chat_openai_instance.get_num_tokens = MagicMock(return_value=0) token_override = getattr(context, "openai_token_count", None) if token_override is not None: mock_chat_openai_instance.get_num_tokens.return_value = token_override mock_chat_openai_class.return_value = mock_chat_openai_instance organization = getattr(context, "openai_provider_org", None) extra_kwargs = dict(getattr(context, "openai_provider_kwargs", {})) context.provider = OpenAIChatProvider( api_key=api_key, model=model_id, organization=organization, **extra_kwargs, ) context.chat_openai_class = mock_chat_openai_class context.chat_openai_instance = mock_chat_openai_instance @when("I attempt to create an OpenAI chat provider without an API key") def step_openai_provider_without_api_key(context): context.openai_provider_error = None try: OpenAIChatProvider(api_key="", model="gpt-4o-mini") except Exception as exc: # pragma: no cover - defensive logging only context.openai_provider_error = exc @when("I request plan generation from the OpenAI 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_openai_call = context.chat_openai_class.call_args context.plan_generation_graph_call = context.plan_generation_graph_class.call_args @when("I stream plan generation from the OpenAI 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_openai_call = context.chat_openai_class.call_args context.plan_generation_graph_call = context.plan_generation_graph_class.call_args @then( 'the OpenAI provider should construct ChatOpenAI with api key "{api_key}" and model "{model_id}"' ) def step_assert_chat_openai_constructor(context, api_key, model_id): assert context.chat_openai_call is not None, "ChatOpenAI should have been called" call_args, call_kwargs = context.chat_openai_call assert call_args == (), "ChatOpenAI should be called with keyword arguments" assert call_kwargs["api_key"] == api_key assert call_kwargs["model"] == model_id assert context.plan_generation_graph_call is not None, ( "PlanGenerationGraph should receive the ChatOpenAI instance" ) _, graph_kwargs = context.plan_generation_graph_call assert graph_kwargs["llm"] is context.chat_openai_instance @then( 'the OpenAI provider should include organization "{expected_organization}" in the ChatOpenAI call' ) def step_assert_chat_openai_org(context, expected_organization): assert context.chat_openai_call is not None, "ChatOpenAI should have been called" _, call_kwargs = context.chat_openai_call normalized = ( None if expected_organization.lower() == "none" else expected_organization ) assert call_kwargs.get("organization") == normalized @then( 'the OpenAI provider should include kwargs "{kwargs_string}" in the ChatOpenAI call' ) def step_assert_chat_openai_kwargs(context, kwargs_string): assert context.chat_openai_call is not None, "ChatOpenAI should have been called" _, call_kwargs = context.chat_openai_call expected_kwargs = _parse_kwargs_string(kwargs_string) for key, value in expected_kwargs.items(): assert key in call_kwargs, f"Expected {key} in ChatOpenAI kwargs" assert call_kwargs[key] == value, ( f"Expected ChatOpenAI kwargs[{key!r}] to equal {value!r}" ) @then( 'the OpenAI 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 OpenAI 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 OpenAI 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 OpenAI provider response should include {count:d} generated change for "{file_path}"' ) def step_assert_generated_change(context, count, file_path): assert context.response is not None, "Provider response should exist" assert len(context.response.changes) == count, ( f"Expected {count} changes, got {len(context.response.changes)}" ) assert any( getattr(change, "file_path", None) == file_path for change in context.response.changes ), f"Expected change for {file_path}" @then("the OpenAI provider response token count should equal {expected:d}") def step_assert_token_count(context, expected): assert context.response is not None assert context.response.token_count == expected @then('the OpenAI 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 OpenAI 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 OpenAI 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 OpenAI provider creation should fail with error "{message}"') def step_assert_openai_provider_error(context, message): error = getattr(context, "openai_provider_error", None) assert error is not None, "Expected the provider to raise an error" assert isinstance(error, ValueError) assert str(error) == message