use crate::message::AgentMessage;
use crate::types::{
AfterToolCallContext, AfterToolCallResult, AgentLoopTurnUpdate, BeforeToolCallContext,
BeforeToolCallResult, ShouldStopAfterTurnContext, ToolExecutionMode,
};
use futures::future::BoxFuture;
use rpi_ai::provider::{CacheRetention, SimpleStreamOptions};
use rpi_ai::types::{Message, ThinkingLevel};
use rpi_ai::Model;
use std::sync::Arc;
use std::time::Duration;
use tokio_util::sync::CancellationToken;
pub type ConvertToLlm =
Arc<dyn Fn(Vec<AgentMessage>) -> BoxFuture<'static, Vec<Message>> + Send + Sync>;
pub type TransformContext = Arc<
dyn Fn(Vec<AgentMessage>, CancellationToken) -> BoxFuture<'static, Vec<AgentMessage>>
+ Send
+ Sync,
>;
pub type GetApiKey = Arc<dyn Fn(&str) -> BoxFuture<'static, Option<String>> + Send + Sync>;
pub type ShouldStopAfterTurn =
Arc<dyn Fn(ShouldStopAfterTurnContext<'_>) -> BoxFuture<'static, bool> + Send + Sync>;
pub type PrepareNextTurn = Arc<
dyn Fn(ShouldStopAfterTurnContext<'_>) -> BoxFuture<'static, Option<AgentLoopTurnUpdate>>
+ Send
+ Sync,
>;
pub type GetSteeringMessages = Arc<dyn Fn() -> BoxFuture<'static, Vec<AgentMessage>> + Send + Sync>;
pub type GetFollowUpMessages = Arc<dyn Fn() -> BoxFuture<'static, Vec<AgentMessage>> + Send + Sync>;
pub type BeforeToolCall = Arc<
dyn Fn(
BeforeToolCallContext<'_>,
CancellationToken,
) -> BoxFuture<'static, Option<BeforeToolCallResult>>
+ Send
+ Sync,
>;
pub type AfterToolCall = Arc<
dyn Fn(
AfterToolCallContext<'_>,
CancellationToken,
) -> BoxFuture<'static, Option<AfterToolCallResult>>
+ Send
+ Sync,
>;
#[derive(Clone)]
pub struct AgentLoopConfig {
pub model: Model,
pub convert_to_llm: ConvertToLlm,
pub transform_context: Option<TransformContext>,
pub get_api_key: Option<GetApiKey>,
pub should_stop_after_turn: Option<ShouldStopAfterTurn>,
pub prepare_next_turn: Option<PrepareNextTurn>,
pub get_steering_messages: Option<GetSteeringMessages>,
pub get_follow_up_messages: Option<GetFollowUpMessages>,
pub before_tool_call: Option<BeforeToolCall>,
pub after_tool_call: Option<AfterToolCall>,
pub tool_execution: ToolExecutionMode,
pub thinking_level: ThinkingLevel,
pub api_key: Option<String>,
pub timeout: Option<Duration>,
pub max_retries: Option<u32>,
pub max_retry_delay: Option<Duration>,
pub cache_retention: CacheRetention,
pub session_id: Option<String>,
pub signal: CancellationToken,
}
impl std::fmt::Debug for AgentLoopConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AgentLoopConfig")
.field("model", &self.model)
.field("tool_execution", &self.tool_execution)
.field("thinking_level", &self.thinking_level)
.field("cache_retention", &self.cache_retention)
.field("session_id", &self.session_id)
.field("transform_context", &self.transform_context.is_some())
.field("get_api_key", &self.get_api_key.is_some())
.field(
"should_stop_after_turn",
&self.should_stop_after_turn.is_some(),
)
.field("prepare_next_turn", &self.prepare_next_turn.is_some())
.field(
"get_steering_messages",
&self.get_steering_messages.is_some(),
)
.field(
"get_follow_up_messages",
&self.get_follow_up_messages.is_some(),
)
.field("before_tool_call", &self.before_tool_call.is_some())
.field("after_tool_call", &self.after_tool_call.is_some())
.finish()
}
}
impl AgentLoopConfig {
pub fn to_stream_options(&self, api_key: Option<String>) -> SimpleStreamOptions {
let mut opts = SimpleStreamOptions {
api_key,
timeout: self.timeout,
max_retries: self.max_retries,
max_retry_delay: self.max_retry_delay,
headers: None,
metadata: None,
cache_retention: self.cache_retention,
session_id: self.session_id.clone(),
signal: self.signal.clone(),
..SimpleStreamOptions::default()
};
match self.thinking_level {
ThinkingLevel::Off => opts.reasoning = None,
other => opts.reasoning = Some(other),
}
opts
}
}
pub fn default_convert_to_llm(messages: Vec<AgentMessage>) -> Vec<Message> {
messages
.into_iter()
.filter_map(|m| match m {
AgentMessage::User(u) => Some(Message::User(u)),
AgentMessage::Assistant(a) => Some(Message::Assistant(a)),
AgentMessage::ToolResult(t) => Some(Message::ToolResult(t)),
AgentMessage::Custom(_) => None,
})
.collect()
}
pub fn default_convert_to_llm_fn() -> ConvertToLlm {
Arc::new(|messages: Vec<AgentMessage>| {
let out = default_convert_to_llm(messages);
Box::pin(async move { out })
})
}