diff --git a/src/cai/agents/patterns/redteam_codeagent_as_tool.py b/src/cai/agents/patterns/redteam_codeagent_as_tool.py deleted file mode 100644 index fdbb195e..00000000 --- a/src/cai/agents/patterns/redteam_codeagent_as_tool.py +++ /dev/null @@ -1,63 +0,0 @@ -""" -Implementation of a Red Team Agent with Code Execution Agent as Tool Pattern - -This module establishes a specialized pattern where a Code Execution Agent -serves as a tool for the Red Team Agent. This allows the Red Team Agent to -delegate code execution tasks to a dedicated agent, enhancing security -assessment capabilities through specialized code execution. -""" -import os -from dotenv import load_dotenv -from openai import AsyncOpenAI - -from cai.sdk.agents import Agent, OpenAIChatCompletionsModel -from cai.agents.red_teamer import redteam_agent -from cai.tools.reconnaissance.exec_code import execute_code -from cai.util import load_prompt_template - - -# Load environment variables -load_dotenv() -model_name = os.getenv("CAI_MODEL", "qwen2.5:14b") - -# Create a specialized code execution agent with only the execute_code tool -code_execution_agent = Agent( - name="Code Execution Agent", - description="Specialized agent for executing code in security assessments", - instructions=( - "You are a specialized code execution agent that helps with security " - "assessments. Execute the task and then transfer to red team" - ), - tools=[execute_code], - handoffs=[], - model=OpenAIChatCompletionsModel( - model=model_name, - openai_client=AsyncOpenAI(), - ), -) - -# Create a handoff for the red team agent -from cai.sdk.agents import handoff - -redteam_handoff = handoff( - agent=redteam_agent, - tool_description_override="Transfer to Red Team Agent for security assessment and exploitation tasks" -) - -# Add the handoff to the code execution agent -code_execution_agent.handoffs.append(redteam_handoff) - -# Register the code execution agent as a tool for the red team agent -redteam_agent.tools.append( - code_execution_agent.as_tool( - tool_name="code_execution", - tool_description=( - "Use this tool when you need to execute code" - ), - ) -) -if execute_code in redteam_agent.tools: - redteam_agent.tools.remove(execute_code) -# Export the enhanced red team agent as the pattern -redteam_with_code_agent_pattern = redteam_agent -redteam_with_code_agent_pattern.pattern = "agent_as_tool" diff --git a/src/cai/sdk/agents/function_schema.py b/src/cai/sdk/agents/function_schema.py index 681affce..193c5c6d 100644 --- a/src/cai/sdk/agents/function_schema.py +++ b/src/cai/sdk/agents/function_schema.py @@ -36,7 +36,6 @@ class FuncSchema: strict_json_schema: bool = True """Whether the JSON schema is in strict mode. We **strongly** recommend setting this to True, as it increases the likelihood of correct JSON input.""" - def to_call_args(self, data: BaseModel) -> tuple[list[Any], dict[str, Any]]: """ Converts validated data from the Pydantic model into (args, kwargs), suitable for calling @@ -51,6 +50,10 @@ class FuncSchema: # If the function takes a RunContextWrapper and this is the first parameter, skip it. if self.takes_context and idx == 0: continue + + # Skip parameters named 'ctf' or 'CTF' + if name.lower() == 'ctf': + continue value = getattr(data, name, None) if param.kind == param.VAR_POSITIONAL: diff --git a/src/cai/sdk/agents/model_settings.py b/src/cai/sdk/agents/model_settings.py index 2b0885ab..2d413576 100644 --- a/src/cai/sdk/agents/model_settings.py +++ b/src/cai/sdk/agents/model_settings.py @@ -43,6 +43,10 @@ class ModelSettings: store: bool | None = None """Whether to store the generated model response for later retrieval. Defaults to True if not provided.""" + + agent_model: str | None = None + """The model from the Agent class. If set, this will override the model provided + to the OpenAIChatCompletionsModel during initialization.""" def resolve(self, override: ModelSettings | None) -> ModelSettings: """Produce a new ModelSettings by overlaying any non-None values from the diff --git a/src/cai/sdk/agents/models/openai_chatcompletions.py b/src/cai/sdk/agents/models/openai_chatcompletions.py index b2e72417..2a023b16 100644 --- a/src/cai/sdk/agents/models/openai_chatcompletions.py +++ b/src/cai/sdk/agents/models/openai_chatcompletions.py @@ -1485,9 +1485,16 @@ class OpenAIChatCompletionsModel(Model): # Match the behavior of Responses where store is True when not given store = model_settings.store if model_settings.store is not None else True + # Check if we should use the agent's model instead of self.model + # This prioritizes the model from Agent when available + agent_model = None + if hasattr(model_settings, 'agent_model') and model_settings.agent_model: + agent_model = model_settings.agent_model + logger.debug(f"Using agent model: {agent_model} instead of {self.model}") + # Prepare kwargs for the API call kwargs = { - "model": self.model, + "model": agent_model if agent_model else self.model, "messages": converted_messages, "tools": converted_tools or NOT_GIVEN, "temperature": self._non_null_or_not_given(model_settings.temperature), @@ -1505,7 +1512,7 @@ class OpenAIChatCompletionsModel(Model): } # Determine provider based on model string - model_str = str(self.model).lower() + model_str = str(kwargs["model"]).lower() # Provider-specific adjustments if "/" in model_str: @@ -1762,13 +1769,10 @@ class OpenAIChatCompletionsModel(Model): parallel_tool_calls: bool ) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]: """Handle standard LiteLLM API calls for OpenAI and compatible models.""" - # Make sure model is the first parameter - model = kwargs.pop("model", self.model) - if stream: # Standard LiteLLM handling for streaming - ret = litellm.completion(model=model, **kwargs) - stream_obj = await litellm.acompletion(model=model, **kwargs) + ret = litellm.completion(**kwargs) + stream_obj = await litellm.acompletion(**kwargs) response = Response( id=FAKE_RESPONSES_ID, @@ -1785,7 +1789,7 @@ class OpenAIChatCompletionsModel(Model): return response, stream_obj else: # Standard OpenAI handling for non-streaming - ret = litellm.completion(model=model, **kwargs) + ret = litellm.completion(**kwargs) return ret async def _fetch_response_litellm_ollama( @@ -1797,11 +1801,9 @@ class OpenAIChatCompletionsModel(Model): parallel_tool_calls: bool, provider="ollama" ) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]: - # Extract the model first to ensure it's the first parameter - model = kwargs.get("model", self.model) - # Extract only supported parameters for Ollama ollama_supported_params = { + "model": kwargs.get("model", ""), "messages": kwargs.get("messages", []), "stream": kwargs.get("stream", False) } @@ -1826,7 +1828,7 @@ class OpenAIChatCompletionsModel(Model): ollama_kwargs = {k: v for k, v in ollama_supported_params.items() if v is not None} # Check if this is a Qwen model - model_str = str(model).lower() + model_str = str(self.model).lower() is_qwen = "qwen" in model_str api_base = get_ollama_api_base() @@ -1847,21 +1849,21 @@ class OpenAIChatCompletionsModel(Model): tools=[], parallel_tool_calls=parallel_tool_calls or False, ) - # Get streaming response - ensure model is first parameter + # Get streaming response stream_obj = await litellm.acompletion( - model=model, + **ollama_kwargs, api_base=api_base, custom_llm_provider=provider, - **ollama_kwargs ) return response, stream_obj else: - # Get completion response - ensure model is first parameter + + + # Get completion response return litellm.completion( - model=model, + **ollama_kwargs, api_base=api_base, custom_llm_provider=provider, - **ollama_kwargs ) def _intermediate_logs(self): diff --git a/src/cai/sdk/agents/run.py b/src/cai/sdk/agents/run.py index 3bc79a19..c5872fba 100644 --- a/src/cai/sdk/agents/run.py +++ b/src/cai/sdk/agents/run.py @@ -646,6 +646,13 @@ class Runner: model = cls._get_model(agent, run_config) model_settings = agent.model_settings.resolve(run_config.model_settings) model_settings = RunImpl.maybe_reset_tool_choice(agent, tool_use_tracker, model_settings) + + # Ensure agent model is set in model_settings for streaming mode + if not hasattr(model_settings, 'agent_model') or not model_settings.agent_model: + if isinstance(agent.model, str): + model_settings.agent_model = agent.model + elif isinstance(run_config.model, str): + model_settings.agent_model = run_config.model final_response: ModelResponse | None = None @@ -930,6 +937,13 @@ class Runner: model_settings = agent.model_settings.resolve(run_config.model_settings) model_settings = RunImpl.maybe_reset_tool_choice(agent, tool_use_tracker, model_settings) + # Ensure agent model is set in model_settings + if not hasattr(model_settings, 'agent_model') or not model_settings.agent_model: + if isinstance(agent.model, str): + model_settings.agent_model = agent.model + elif isinstance(run_config.model, str): + model_settings.agent_model = run_config.model + new_response = await model.get_response( system_instructions=system_prompt, input=input, @@ -970,14 +984,21 @@ class Runner: @classmethod def _get_model(cls, agent: Agent[Any], run_config: RunConfig) -> Model: model = None + agent_model = None if isinstance(run_config.model, Model): model = run_config.model elif isinstance(run_config.model, str): model = run_config.model_provider.get_model(run_config.model) + agent_model = run_config.model elif isinstance(agent.model, Model): model = agent.model else: model = run_config.model_provider.get_model(agent.model) + agent_model = agent.model + + # Store the original agent model in model_settings for later use + if agent_model and hasattr(agent, 'model_settings'): + agent.model_settings.agent_model = agent_model # Set agent name if the model supports it (for CLI display) if hasattr(model, 'set_agent_name'):