use std::sync::atomic::AtomicBool;
use tokio::sync::mpsc;
use crate::config::AgentConfig;
use crate::llm::{
CompleteChatRetryingParams, LlmCompleteError, LlmRetryingTransportOpts, complete_chat_retrying,
tool_chat_request,
};
use crate::types::{LlmSeedOverride, Message};
pub(crate) struct PerPlanCallModelParams<'a> {
pub llm_backend: &'a (dyn crate::llm::ChatCompletionsBackend + 'static),
pub client: &'a reqwest::Client,
pub api_key: &'a str,
pub cfg: &'a AgentConfig,
pub tools_defs: &'a [crate::types::Tool],
pub messages: &'a [Message],
pub out: Option<&'a mpsc::Sender<String>>,
pub no_stream: bool,
pub cancel: Option<&'a AtomicBool>,
pub temperature_override: Option<f32>,
pub model_override: Option<&'a str>,
pub seed_override: LlmSeedOverride,
pub request_chrome_trace: Option<std::sync::Arc<crate::request_chrome_trace::RequestTurnTrace>>,
pub executor_api_base: Option<&'a str>,
pub executor_api_key: Option<&'a str>,
pub turn_budget: Option<&'a std::sync::Arc<crate::agent::turn_budget::TurnBudgetCounter>>,
pub provider_usage_sink:
Option<&'a std::sync::Arc<std::sync::Mutex<Option<crate::cm_types::Usage>>>>,
}
pub(crate) async fn per_plan_call_model_retrying(
p: PerPlanCallModelParams<'_>,
) -> Result<(Message, String), LlmCompleteError> {
let PerPlanCallModelParams {
llm_backend,
client,
api_key,
cfg,
tools_defs,
messages,
out,
no_stream,
cancel,
temperature_override,
model_override,
seed_override,
request_chrome_trace,
executor_api_base,
executor_api_key,
turn_budget,
provider_usage_sink,
} = p;
let (effective_cfg, effective_api_key) =
if executor_api_base.is_some() || executor_api_key.is_some() {
let mut c = (*cfg).clone();
let mut key = api_key.to_string();
if let Some(base) = executor_api_base {
c.llm.api_base = base.to_string();
}
if let Some(key_override) = executor_api_key {
key = key_override.to_string();
}
(std::sync::Arc::new(c), key)
} else {
(std::sync::Arc::new(cfg.clone()), api_key.to_string())
};
let llm_cfg = crate::cm_types::llm_config::LlmConfig {
llm: effective_cfg.llm.clone(),
sampling: effective_cfg.llm_sampling.clone(),
vendor_flags: effective_cfg.llm_vendor_flags.clone(),
http_retry: effective_cfg.llm_http_retry.clone(),
};
let req = tool_chat_request(
&llm_cfg,
messages,
tools_defs,
temperature_override,
model_override,
seed_override,
);
let cc = CompleteChatRetryingParams::new(
llm_backend,
client,
&effective_api_key,
&effective_cfg,
LlmRetryingTransportOpts {
out,
no_stream,
cancel,
},
request_chrome_trace,
model_override,
)
.with_turn_budget(turn_budget)
.with_provider_usage_sink(provider_usage_sink);
let (msg, finish_reason) = complete_chat_retrying(&cc, &req).await?;
Ok((msg, finish_reason))
}