diff --git a/src/cai/agents/flag_discriminator.py b/src/cai/agents/flag_discriminator.py index 6510a7ef..b638448f 100644 --- a/src/cai/agents/flag_discriminator.py +++ b/src/cai/agents/flag_discriminator.py @@ -2,12 +2,11 @@ CTF Flag Discriminator Agent with test """ import os -from cai.sdk.agents import Agent +from cai.sdk.agents import Agent, OpenAIChatCompletionsModel from openai import AsyncOpenAI model = os.getenv('CAI_MODEL', "qwen2.5:14b") flag_discriminator = Agent( - model="gpt-4o" if os.getenv('CAI_MODEL') == "o3-mini" else model, name="Flag discriminator", description="Agent focused on extracting the flag from the output", instructions="""INSTRUCTIONS: @@ -15,7 +14,11 @@ flag_discriminator = Agent( 2. ONLY return the flag, nothing else. 3. Flags can be in any format (not always CTF{...}). 4. If you do not find a flag, call `ctf_agent` to continue investigating. - """ + """, + model=OpenAIChatCompletionsModel( + model="gpt-4o" if os.getenv('CAI_MODEL') == "o3-mini" else model, + openai_client=AsyncOpenAI(), + ) ) # Transfer Function diff --git a/src/cai/agents/patterns/red_team.py b/src/cai/agents/patterns/red_team.py index 01d56b52..00b408ad 100644 --- a/src/cai/agents/patterns/red_team.py +++ b/src/cai/agents/patterns/red_team.py @@ -37,10 +37,10 @@ def thought_agent_handoff(ctf=None): # pylint: disable=unused-argument return thought_agent -# Register handoff functions to enable inter-agent communication pathways -redteam_agent.functions.append(transfer_to_dns_agent) -dns_smtp_agent.functions.append(redteam_agent_handoff) -thought_agent.functions.append(redteam_agent_handoff) +# Register handoff to enable inter-agent communication pathways +redteam_agent.handoffs.append(transfer_to_dns_agent) +dns_smtp_agent.handoffs.append(redteam_agent_handoff) +thought_agent.handoffs.append(redteam_agent_handoff) # Initialize the swarm pattern with the thought agent as the entry point redteam_swarm_pattern = thought_agent diff --git a/tests/agents/test_agent_inference.py b/tests/agents/test_agent_inference.py index e3b9752c..43b690d2 100644 --- a/tests/agents/test_agent_inference.py +++ b/tests/agents/test_agent_inference.py @@ -47,7 +47,6 @@ async def test_one_tool_agent_inference(): Non-streaming inference test for the one_tool_agent. """ prompt = "use nmap to scan 192.168.2.0/24" - # Run the agent in non-streaming mode result = await Runner.run(get_agent_by_name("one_tool_agent"), prompt) final_output = result.final_output or "" assert final_output, "Expected non-empty final output" @@ -65,3 +64,30 @@ async def test_red_team_agent_inference(): final_output = result.final_output or "" assert final_output, "Expected non-empty final output" assert "example.com" in final_output.lower(), f"Expected 'example.com' in output, got: {final_output}" + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_flag_discriminator_inference(): + """ + Non-streaming inference test for the one_tool_agent. + """ + prompt = "Hello! Can you find tell me which is the flag in this string: 'Hi there, your reward flag{1234}" + result = await Runner.run(get_agent_by_name("flag_discriminator"), prompt) + final_output = result.final_output or "" + assert final_output, "Expected non-empty final output" + assert "flag{1234}" in final_output.lower(), f"Expected 'flag{1234}' in output, got: {final_output}" + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_thought_agent_inference(): + """ + Non-streaming inference test for the one_tool_agent. + """ + prompt = """The phases of the cybersecurity kill chain are: + Reconnaissance, Exploitation, Lateral Movement, Data Exfiltration, and Command and Control (C2). + In which phase of the kill chain does this example fall: + Gathering initial intelligence about the target using OSINT techniques?""" + result = await Runner.run(get_agent_by_name("thought_agent"), prompt) + final_output = result.final_output or "" + assert final_output, "Expected non-empty final output" + assert "reconnaissance" in final_output.lower(), f"Expected 'reconnaissance' in output, got: {final_output}" \ No newline at end of file