Files
cleveragents-core/features/steps/openai_provider_steps.py
T
HAL9000 9b6bedb463
CI / lint (pull_request) Failing after 43s
CI / helm (pull_request) Successful in 34s
CI / build (pull_request) Successful in 39s
CI / push-validation (pull_request) Successful in 27s
CI / quality (pull_request) Successful in 59s
CI / typecheck (pull_request) Successful in 1m20s
CI / security (pull_request) Successful in 1m49s
CI / unit_tests (pull_request) Successful in 5m57s
CI / coverage (pull_request) Has been skipped
CI / docker (pull_request) Has been skipped
CI / integration_tests (pull_request) Successful in 11m25s
CI / status-check (pull_request) Failing after 4s
fix(tests): resolve AmbiguousStep errors in provider BDD step files
Extract shared Behave step definitions for LLM provider tests into
provider_shared_steps.py to eliminate duplicate @given registrations
that caused AmbiguousStep errors across anthropic, google, and openai
provider step files. Add missing step definitions to
consolidated_ai_models_providers_steps.py for consolidated feature
scenarios. Add @given decorator alongside @when for provider creation
steps used as Given steps in feature files.
2026-06-04 18:44:59 -04:00

295 lines
11 KiB
Python

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.openai_provider import OpenAIChatProvider
# Shared step definitions (I have sample provider domain inputs, plan generation graph steps)
# are registered by features/steps/provider_shared_steps.py which Behave loads automatically.
# Helper functions are duplicated here to avoid cross-module import issues with Behave's exec_file.
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("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('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