diff --git a/.gitignore b/.gitignore index bee4b907..e66cd370 100644 --- a/.gitignore +++ b/.gitignore @@ -146,4 +146,4 @@ cython_debug/ # CAI files .cai/ .vscode/ - +cai_env/ diff --git a/tests/agents/test_agent_one_tool.py b/tests/agents/test_agent_one_tool.py new file mode 100644 index 00000000..9290cef8 --- /dev/null +++ b/tests/agents/test_agent_one_tool.py @@ -0,0 +1,33 @@ +import pytest +from tests.fake_model import FakeModel +from tests.test_responses import ( + get_text_message, + get_function_tool_call, + get_function_tool, +) +from cai.sdk.agents import Runner +from cai.agents.one_tool import transfer_to_one_tool_agent + +@pytest.mark.asyncio +async def test_ctf_agent_executes_linux_command(): + model = FakeModel() + agent = transfer_to_one_tool_agent() + agent.model = model + model.add_multiple_turn_outputs( + [ + [ + get_text_message("executing comando..."), + get_function_tool_call("generic_linux_command", '{"command": "ls"}') + ], + [ + get_text_message("result of the command: flag{12345}") + ] + ] + ) + + result = await Runner.run(agent, input="List files") + + assert result.final_output == "result of the command: flag{12345}" + assert len(result.raw_responses) == 2 + + assert any("generic_linux_command" in str(item) for item in result.to_input_list()) diff --git a/tests/agents/test_agent_tracing.py b/tests/agents/test_agent_tracing.py index 4f6f5c69..446a99e7 100644 --- a/tests/agents/test_agent_tracing.py +++ b/tests/agents/test_agent_tracing.py @@ -7,9 +7,9 @@ from inline_snapshot import snapshot from cai.sdk.agents import Agent, RunConfig, Runner, trace -from .fake_model import FakeModel -from .test_responses import get_text_message -from .testing_processor import assert_no_traces, fetch_normalized_spans +from tests.fake_model import FakeModel +from tests.test_responses import get_text_message +from tests.testing_processor import assert_no_traces, fetch_normalized_spans @pytest.mark.asyncio diff --git a/tests/agents/test_function_tool.py b/tests/tools/test_function_tool.py similarity index 100% rename from tests/agents/test_function_tool.py rename to tests/tools/test_function_tool.py diff --git a/tests/tools/test_tool_choice_reset.py b/tests/tools/test_tool_choice_reset.py index 9244ff98..220ca4c2 100644 --- a/tests/tools/test_tool_choice_reset.py +++ b/tests/tools/test_tool_choice_reset.py @@ -3,8 +3,8 @@ import pytest from cai.sdk.agents import Agent, ModelSettings, Runner from cai.sdk.agents._run_impl import AgentToolUseTracker, RunImpl -from .fake_model import FakeModel -from .test_responses import get_function_tool, get_function_tool_call, get_text_message +from tests.fake_model import FakeModel +from tests.test_responses import get_function_tool, get_function_tool_call, get_text_message class TestToolChoiceReset: diff --git a/tests/tools/test_tool_use_behavior.py b/tests/tools/test_tool_use_behavior.py index 6d1935f3..adf25946 100644 --- a/tests/tools/test_tool_use_behavior.py +++ b/tests/tools/test_tool_use_behavior.py @@ -18,7 +18,7 @@ from cai.sdk.agents import ( ) from cai.sdk.agents._run_impl import RunImpl -from .test_responses import get_function_tool +from tests.test_responses import get_function_tool def _make_function_tool_result( diff --git a/tests/tracing/test_tracing.py b/tests/tracing/test_tracing.py index d1e96044..a796efda 100644 --- a/tests/tracing/test_tracing.py +++ b/tests/tracing/test_tracing.py @@ -18,7 +18,7 @@ from cai.sdk.agents.tracing import ( ) from cai.sdk.agents.tracing.spans import SpanError -from .testing_processor import ( +from tests.testing_processor import ( SPAN_PROCESSOR_TESTING, assert_no_traces, fetch_events, diff --git a/tests/tracing/test_tracing_errors.py b/tests/tracing/test_tracing_errors.py index db1fa554..584b77b2 100644 --- a/tests/tracing/test_tracing_errors.py +++ b/tests/tracing/test_tracing_errors.py @@ -19,15 +19,15 @@ from cai.sdk.agents import ( TResponseInputItem, ) -from .fake_model import FakeModel -from .test_responses import ( +from tests.fake_model import FakeModel +from tests.test_responses import ( get_final_output_message, get_function_tool, get_function_tool_call, get_handoff_tool_call, get_text_message, ) -from .testing_processor import fetch_normalized_spans +from tests.testing_processor import fetch_normalized_spans @pytest.mark.asyncio diff --git a/tests/tracing/test_tracing_errors_streamed.py b/tests/tracing/test_tracing_errors_streamed.py index fa0a1e4e..088d86df 100644 --- a/tests/tracing/test_tracing_errors_streamed.py +++ b/tests/tracing/test_tracing_errors_streamed.py @@ -22,15 +22,15 @@ from cai.sdk.agents import ( TResponseInputItem, ) -from .fake_model import FakeModel -from .test_responses import ( +from tests.fake_model import FakeModel +from tests.test_responses import ( get_final_output_message, get_function_tool, get_function_tool_call, get_handoff_tool_call, get_text_message, ) -from .testing_processor import fetch_normalized_spans +from tests.testing_processor import fetch_normalized_spans @pytest.mark.asyncio