use std::sync::mpsc;
#[cfg(feature = "ssh")]
pub const PROVIDER_CREDENTIAL_ENV_VARS: &[&str] = &[
"MARS_LLM_KEY",
"ARES_LLM_KEY",
"AWS_BEARER_TOKEN_BEDROCK",
"AZURE_OPENAI_API_KEY",
"ANTHROPIC_API_KEY",
"OPENAI_API_KEY",
"GROQ_API_KEY",
"GEMINI_API_KEY",
"GOOGLE_API_KEY",
];
#[derive(Clone, Debug, PartialEq)]
pub enum AgentDirective {
Run(String),
Type(String),
Open(String),
Need(NeedKind),
}
#[derive(Clone, Debug, PartialEq)]
pub enum NeedKind {
Scrollback,
Tab(String),
}
pub enum AgentEvent {
Answer {
text: String,
directive: Option<AgentDirective>,
},
AnswerStart,
AnswerDelta { text: String },
AutoName { tab_id: usize, name: String },
SessionName { name: String },
WatchSummary { term_id: usize, verdict: String },
SurfaceSummary { term_id: usize, text: String },
Mission { text: String },
BgDone,
ShellTranslation { command: String, call_id: u64 },
ShiftDelta { text: String },
ShiftDone,
Goals { goals: Vec<String> },
Error(String),
}
fn kebab(text: &str) -> String {
let s: String = text
.trim()
.to_lowercase()
.chars()
.map(|c| if c.is_alphanumeric() { c } else { '-' })
.collect::<String>()
.split('-')
.filter(|s| !s.is_empty())
.collect::<Vec<_>>()
.join("-");
s.chars().take(16).collect()
}
fn match_directive(line: &str) -> Option<AgentDirective> {
let l = line
.trim()
.trim_start_matches(['-', '*', '>', ' '])
.trim_matches('`')
.trim_matches('*')
.trim();
if let Some(rest) = l.strip_prefix("RUN:") {
if let Some(name) = rest.trim().trim_matches('`').split_whitespace().next() {
let name = name.trim_end_matches(['.', ',', ':']);
if !name.is_empty() {
return Some(AgentDirective::Run(name.to_string()));
}
}
}
if let Some(rest) = l.strip_prefix("NEED:") {
let arg = rest.trim().trim_matches('`').trim();
let low = arg.to_lowercase();
if low.starts_with("scrollback") || low.starts_with("history") {
return Some(AgentDirective::Need(NeedKind::Scrollback));
}
if let Some(tab) = low.strip_prefix("tab") {
let name = tab.trim().to_string();
if !name.is_empty() {
return Some(AgentDirective::Need(NeedKind::Tab(name)));
}
}
}
for (tag, make) in [
("TYPE:", AgentDirective::Type as fn(String) -> AgentDirective),
("OPEN:", AgentDirective::Open as fn(String) -> AgentDirective),
] {
if let Some(rest) = l.strip_prefix(tag) {
let arg = rest.trim().trim_matches('`').trim().to_string();
if !arg.is_empty() {
return Some(make(arg));
}
}
}
None
}
pub fn parse_directive(text: &str) -> (String, Option<AgentDirective>) {
let lines: Vec<&str> = text.lines().collect();
let mut hit: Option<usize> = None;
for (i, line) in lines.iter().enumerate().rev().take(4) {
if line.trim().is_empty() {
continue;
}
if match_directive(line).is_some() {
hit = Some(i);
break;
}
}
match hit {
Some(i) => {
let directive = match_directive(lines[i]);
let display = lines[..i].join("\n").trim_end().to_string();
(display, directive)
}
None => (text.to_string(), None),
}
}
#[derive(Clone)]
pub struct AgentConfig {
pub url: String,
pub key: String,
pub model: String,
pub provider: &'static str,
pub max_tokens: u32,
pub temperature: f64,
pub broker_sock: Option<String>,
}
fn env_var(name: &str) -> Result<String, std::env::VarError> {
std::env::var(format!("MARS_{name}")).or_else(|_| std::env::var(format!("ARES_{name}")))
}
#[derive(Debug)]
pub struct RateLimited(pub String);
impl std::fmt::Display for RateLimited {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.0.fmt(f)
}
}
impl std::error::Error for RateLimited {}
#[derive(Debug)]
pub struct ModelUnavailable(pub String);
impl std::fmt::Display for ModelUnavailable {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.0.fmt(f)
}
}
impl std::error::Error for ModelUnavailable {}
pub fn is_retired_model(code: u16, api_msg: Option<&str>) -> bool {
code == 404
|| api_msg
.map(|m| {
let l = m.to_lowercase();
l.contains("does not exist")
|| l.contains("do not have access")
|| l.contains("decommission")
|| l.contains("has been deprecated")
|| l.contains("model_not_found")
|| l.contains("no longer supported")
})
.unwrap_or(false)
}
type Provider = (String, &'static str, String, String);
fn provider_chain() -> Vec<Provider> {
let mut chain: Vec<Provider> = Vec::new();
let p = |k: String, tag: &'static str, url: &str, model: &str| {
(k, tag, url.to_string(), model.to_string())
};
if let Ok(k) = env_var("LLM_KEY") {
chain.push(p(k, "custom", "https://api.groq.com/openai/v1", "llama-3.1-8b-instant"));
}
if let Ok(k) = std::env::var("AWS_BEARER_TOKEN_BEDROCK") {
let region = env_var("BEDROCK_REGION")
.or_else(|_| std::env::var("AWS_REGION"))
.or_else(|_| std::env::var("AWS_DEFAULT_REGION"))
.unwrap_or_else(|_| "us-east-1".to_string());
chain.push(p(
k,
"bedrock",
&format!("https://bedrock-runtime.{region}.amazonaws.com"),
"us.anthropic.claude-3-5-haiku-20241022-v1:0",
));
}
if let (Ok(k), Ok(endpoint)) =
(std::env::var("AZURE_OPENAI_API_KEY"), std::env::var("AZURE_OPENAI_ENDPOINT"))
{
let deployment = env_var("AZURE_DEPLOYMENT")
.or_else(|_| env_var("LLM_MODEL"))
.unwrap_or_else(|_| "gpt-4o-mini".to_string());
let version = env_var("AZURE_API_VERSION").unwrap_or_else(|_| "2024-10-21".to_string());
let base = endpoint.trim_end_matches('/');
chain.push((
k,
"azure",
format!("{base}/openai/deployments/{deployment}/chat/completions?api-version={version}"),
deployment,
));
}
if let Ok(k) = std::env::var("ANTHROPIC_API_KEY") {
chain.push(p(k, "anthropic", "https://api.anthropic.com", "claude-haiku-4-5"));
}
if let Ok(k) = std::env::var("OPENAI_API_KEY") {
chain.push(p(k, "openai", "https://api.openai.com/v1", "gpt-4o-mini"));
}
if let Ok(k) = std::env::var("GROQ_API_KEY") {
chain.push(p(k, "groq", "https://api.groq.com/openai/v1", "llama-3.3-70b-versatile"));
}
if let Ok(k) = std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY")) {
chain.push(p(
k,
"gemini",
"https://generativelanguage.googleapis.com/v1beta/openai",
"gemini-3.1-flash-lite",
));
}
chain
}
pub fn rotation_candidates(current_provider: &str) -> Vec<AgentConfig> {
if env_var("LLM_MODEL").is_ok()
|| env_var("LLM_URL").is_ok()
|| current_provider == "custom"
{
return Vec::new();
}
provider_chain()
.into_iter()
.filter(|(_, p, _, _)| *p != current_provider)
.map(|(key, provider, url, model)| AgentConfig {
url: url.to_string(),
key,
model: model.to_string(),
provider,
max_tokens: 512,
temperature: 0.3,
broker_sock: None,
})
.collect()
}
impl AgentConfig {
pub fn from_env() -> Self {
if std::env::var("MARS_LLM_KEY").is_err()
&& std::env::var("ARES_LLM_KEY").is_err()
{
if let Some(sock) = crate::broker::detect_broker_sock() {
return AgentConfig {
url: String::new(),
key: String::new(),
model: env_var("LLM_MODEL").unwrap_or_default(),
provider: "broker",
max_tokens: 512,
temperature: 0.3,
broker_sock: Some(sock),
};
}
}
let (key, provider, default_url, default_model) =
provider_chain().into_iter().next().unwrap_or((
String::new(),
"none",
"https://api.groq.com/openai/v1".to_string(),
"llama-3.1-8b-instant".to_string(),
));
let url = env_var("LLM_URL").unwrap_or_else(|_| default_url.to_string());
let model = env_var("LLM_MODEL").unwrap_or_else(|_| default_model.to_string());
AgentConfig { url, key, model, provider, max_tokens: 512, temperature: 0.3, broker_sock: None }
}
pub fn is_configured(&self) -> bool {
if self.provider == "broker" {
return self
.broker_sock
.as_deref()
.map(|s| {
crate::sys::control::probe(std::path::Path::new(s))
== crate::sys::control::Probe::Live
})
.unwrap_or(false);
}
!self.key.is_empty()
}
}
fn system_prompt(registry: &str, screen: &str) -> String {
crate::prompts::ASK_SYSTEM.replace("{registry}", registry).replace("{screen}", screen)
}
pub fn build_ask_messages(
registry: &str,
screen: &str,
history: &[(String, String)],
question: &str,
) -> Vec<serde_json::Value> {
let mut messages = vec![serde_json::json!({
"role": "system", "content": system_prompt(registry, screen)
})];
if let Some(ctx) = crate::retrieval::docs_context_for(question) {
messages.push(serde_json::json!({ "role": "system", "content": ctx }));
}
if let Some(p) = crate::persona::system_message() {
messages.push(p);
}
let start = history.len().saturating_sub(12);
for (role, content) in &history[start..] {
messages.push(serde_json::json!({ "role": role, "content": content }));
}
messages.push(serde_json::json!({ "role": "user", "content": question }));
messages
}
pub fn ask(
cfg: AgentConfig,
question: String,
registry: String,
screen: String,
history: Vec<(String, String)>,
tx: mpsc::Sender<AgentEvent>,
) {
std::thread::spawn(move || {
let mode = crate::retrieval::MemoryMode::from_env();
let messages = build_ask_messages(®istry, &screen, &history, &question);
let _ = tx.send(AgentEvent::AnswerStart);
let txd = tx.clone();
let mut on_delta = move |d: &str| {
let _ = txd.send(AgentEvent::AnswerDelta { text: d.to_string() });
};
match chat_with_id_streaming(&cfg, messages.clone(), "ask", mode.as_str(), &mut on_delta) {
Ok((text, _call_id)) => {
let (display, directive) = parse_directive(&text);
if let Some(AgentDirective::Run(name)) = &directive {
if crate::palette::Action::from_name(name).is_none() {
if let Some(up) = crate::tiers::model_above(cfg.provider, "ask") {
let cfg_up = AgentConfig { model: up, ..cfg.clone() };
let _ = tx.send(AgentEvent::AnswerStart); if let Ok((text2, _)) = chat_with_id_streaming(
&cfg_up,
messages,
"ask_escalated",
mode.as_str(),
&mut on_delta,
) {
let (display2, directive2) = parse_directive(&text2);
let _ = tx.send(AgentEvent::Answer {
text: display2,
directive: directive2,
});
return;
}
}
}
}
let _ = tx.send(AgentEvent::Answer { text: display, directive });
}
Err(e) => {
let _ = tx.send(AgentEvent::Error(e.to_string()));
}
}
});
}
pub fn auto_name(cfg: AgentConfig, tab_id: usize, screen: String, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let messages = format_task_messages(crate::prompts::AUTO_NAME_SYSTEM, &screen);
if let Ok(text) = chat(&cfg, messages, "auto_name") {
let name = kebab(&text);
if !name.is_empty() {
let _ = tx.send(AgentEvent::AutoName { tab_id, name });
}
}
let _ = tx.send(AgentEvent::BgDone); });
}
pub fn watch_summary(
cfg: AgentConfig,
term_id: usize,
reason: crate::app::WatchReason,
tail: String,
tx: mpsc::Sender<AgentEvent>,
) {
std::thread::spawn(move || {
let messages = build_watch_messages(reason, &tail);
match chat(&cfg, messages, "watch") {
Ok(text) => {
let verdict = text.trim().lines().next().unwrap_or("").trim().to_string();
if !verdict.is_empty() {
let _ = tx.send(AgentEvent::WatchSummary { term_id, verdict });
}
}
Err(e) => {
let _ = tx.send(AgentEvent::WatchSummary {
term_id,
verdict: format!("⚠ watch couldn't summarize — {e}"),
});
}
}
let _ = tx.send(AgentEvent::BgDone); });
}
pub fn summarize_surface(cfg: AgentConfig, term_id: usize, tail: String, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let messages = build_watch_messages(crate::app::WatchReason::Quiet, &tail);
match chat(&cfg, messages, "summarize") {
Ok(text) => {
let line = text.trim().lines().next().unwrap_or("").trim().to_string();
let _ = tx.send(AgentEvent::SurfaceSummary {
term_id,
text: if line.is_empty() { "(nothing to summarize)".to_string() } else { line },
});
}
Err(e) => {
let _ = tx.send(AgentEvent::SurfaceSummary { term_id, text: format!("⚠ couldn't summarize — {e}") });
}
}
let _ = tx.send(AgentEvent::BgDone); });
}
pub fn build_watch_messages(reason: crate::app::WatchReason, tail: &str) -> Vec<serde_json::Value> {
let hint = match reason {
crate::app::WatchReason::Exit => crate::prompts::WATCH_HINT_EXIT,
crate::app::WatchReason::Quiet => crate::prompts::WATCH_HINT_QUIET,
};
let mut messages = vec![serde_json::json!({ "role": "system",
"content": crate::prompts::WATCH_SYSTEM.trim_end().replace("{hint}", hint.trim_end()) })];
if let Some(p) = crate::persona::system_message() {
messages.push(p);
}
messages.push(serde_json::json!({ "role": "user", "content": tail }));
messages
}
pub fn infer_mission(cfg: AgentConfig, snapshots: Vec<String>, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let messages = format_task_messages(crate::prompts::MISSION_SYSTEM, &snapshots.join("\n"));
if let Ok(text) = chat(&cfg, messages, "mission") {
let mission = text.trim().lines().next().unwrap_or("").trim().to_string();
if !mission.is_empty() {
let _ = tx.send(AgentEvent::Mission { text: mission });
}
}
let _ = tx.send(AgentEvent::BgDone);
});
}
pub fn name_session(cfg: AgentConfig, screen: String, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let messages = format_task_messages(crate::prompts::NAME_SESSION_SYSTEM, &screen);
if let Ok(text) = chat(&cfg, messages, "name_session") {
let name = kebab(&text);
if !name.is_empty() {
let _ = tx.send(AgentEvent::SessionName { name });
}
}
let _ = tx.send(AgentEvent::BgDone); });
}
fn is_reasoning_model(model: &str) -> bool {
let m = model.to_lowercase();
["qwen3", "qwq", "deepseek-r1", "-r1", "o1-", "o3", "o4-mini", "thinking", "reasoning"]
.iter()
.any(|p| m.contains(p))
}
pub const TASKS: &[&str] = &[
"ask", "translate", "watch", "mission", "auto_name", "name_session", "shift_brief",
"capture_goals",
];
pub fn parse_goals(text: &str) -> Vec<String> {
text.lines()
.map(|l| {
l.trim()
.trim_start_matches(|c: char| c.is_ascii_digit() || matches!(c, '.' | ')' | '-' | '*' | ' '))
.trim()
.to_string()
})
.filter(|l| !l.is_empty())
.take(3)
.collect()
}
pub fn capture_goals(cfg: AgentConfig, evidence: String, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let system = crate::prompts::CAPTURE_GOALS.trim_end().replace("{evidence}", &evidence);
let messages = format_task_messages(&system, "What am I working on?");
match chat(&cfg, messages, "capture_goals") {
Ok(text) => {
let goals = parse_goals(&text);
if !goals.is_empty() {
let _ = tx.send(AgentEvent::Goals { goals });
}
}
Err(e) => eprintln!(
"[mars] goal capture failed (provider {}, model {}): {e}",
cfg.provider, cfg.model
),
}
let _ = tx.send(AgentEvent::BgDone);
});
}
pub fn shift_brief(
cfg: AgentConfig,
away: String,
mission: String,
prev: String,
evidence: String,
tx: mpsc::Sender<AgentEvent>,
) {
std::thread::spawn(move || {
let system = crate::prompts::SHIFT_BRIEF
.trim_end()
.replace("{away}", &away)
.replace("{mission}", if mission.is_empty() { "(none inferred)" } else { &mission })
.replace("{prev}", if prev.is_empty() { "(this is the first briefing)" } else { &prev })
.replace("{evidence}", &evidence);
let mut messages = vec![serde_json::json!({ "role": "system", "content": system })];
if let Some(p) = crate::persona::system_message() {
messages.push(p); }
messages.push(serde_json::json!({ "role": "user", "content": "Report." }));
let streamed = std::cell::Cell::new(false);
let mut on_delta = |d: &str| {
streamed.set(true);
let _ = tx.send(AgentEvent::ShiftDelta { text: d.to_string() });
};
match chat_with_id_streaming(&cfg, messages, "shift_brief", "n/a", &mut on_delta) {
Ok((text, _)) if !streamed.get() && !text.trim().is_empty() => {
let _ = tx.send(AgentEvent::ShiftDelta { text });
}
Err(e) => {
eprintln!(
"[mars] shift_brief enrichment failed (provider {}, model {}): {e}",
cfg.provider, cfg.model
);
}
_ => {} }
let _ = tx.send(AgentEvent::ShiftDone);
let _ = tx.send(AgentEvent::BgDone);
});
}
pub fn build_translate_messages(
reasoning_cap: &str,
examples_block: &str,
request: &str,
screen: &str,
) -> Vec<serde_json::Value> {
let system = crate::prompts::TRANSLATE_SYSTEM
.trim_end()
.replace("{reasoning_cap}", reasoning_cap)
.replace("{examples_block}", examples_block);
vec![
serde_json::json!({ "role": "system", "content": system }),
serde_json::json!({ "role": "user", "content": format!("SCREEN:\n{screen}\n\nREQUEST: {request}") }),
]
}
pub fn format_task_messages(system: &str, user: &str) -> Vec<serde_json::Value> {
vec![
serde_json::json!({ "role": "system", "content": system.trim_end() }),
serde_json::json!({ "role": "user", "content": user }),
]
}
pub fn translate_once(cfg: &AgentConfig, request: &str, screen: &str) -> anyhow::Result<(String, u64)> {
let mode = crate::retrieval::MemoryMode::from_env();
let examples = crate::retrieval::fewshot_for(request);
let reasoning_cap = if is_reasoning_model(&cfg.model) {
format!(" {}", crate::prompts::TRANSLATE_REASONING_CAP.trim())
} else {
String::new()
};
let examples_block = if examples.is_empty() {
String::new()
} else {
format!(
"\n\n{}",
crate::prompts::TRANSLATE_EXAMPLES.trim_end().replace("{examples}", &examples)
)
};
let messages = build_translate_messages(&reasoning_cap, &examples_block, request, screen);
let (text, call_id) = chat_with_id(cfg, messages, "translate", mode.as_str())?;
let command = text
.trim()
.trim_matches('`')
.lines()
.find(|l| !l.trim().is_empty())
.unwrap_or("")
.trim()
.trim_start_matches("$ ")
.to_string();
Ok((command, call_id))
}
pub fn translate_shell(cfg: AgentConfig, request: String, screen: String, tx: mpsc::Sender<AgentEvent>) {
std::thread::spawn(move || {
let ev = match translate_once(&cfg, &request, &screen) {
Ok((command, _)) if command.is_empty() => {
AgentEvent::Error("couldn't translate that — rephrase and retry".into())
}
Ok((command, call_id)) => AgentEvent::ShellTranslation { command, call_id },
Err(e) => AgentEvent::Error(e.to_string()),
};
let _ = tx.send(ev);
});
}
pub fn retry_secs(msg: &str) -> Option<u64> {
let after = msg.split("retry in ").nth(1)?;
let num: String = after.chars().take_while(|c| c.is_ascii_digit() || *c == '.').collect();
num.parse::<f64>().ok().map(|s| s.ceil() as u64)
}
fn strip_reasoning(text: &str) -> String {
let mut out = text.to_string();
while let (Some(a), Some(b)) = (out.find("<think>"), out.find("</think>")) {
if a < b {
out.replace_range(a..b + "</think>".len(), "");
} else {
break;
}
}
if let Some(a) = out.find("<think>") {
out.truncate(a);
}
out.trim().to_string()
}
type DeltaSink<'a> = Option<&'a mut dyn FnMut(&str)>;
fn reborrow<'b>(sink: &'b mut DeltaSink<'_>) -> DeltaSink<'b> {
match sink {
Some(s) => Some(&mut **s),
None => None,
}
}
pub(crate) fn stream_visible(raw: &str) -> String {
let s = strip_reasoning(raw);
let tag = "<think>";
let mut cut = s.len();
for k in (1..tag.len()).rev() {
if s.ends_with(&tag[..k]) {
cut = s.len() - k;
break;
}
}
s[..cut].trim_end().to_string()
}
pub fn chat_with_id(
cfg: &AgentConfig,
messages: Vec<serde_json::Value>,
task: &str,
retrieval: &str,
) -> anyhow::Result<(String, u64)> {
chat_inner(cfg, messages, task, retrieval, None)
}
pub fn chat_with_id_streaming(
cfg: &AgentConfig,
messages: Vec<serde_json::Value>,
task: &str,
retrieval: &str,
on_delta: &mut dyn FnMut(&str),
) -> anyhow::Result<(String, u64)> {
chat_inner(cfg, messages, task, retrieval, Some(on_delta))
}
fn chat_inner(
cfg: &AgentConfig,
messages: Vec<serde_json::Value>,
task: &str,
retrieval: &str,
mut sink: DeltaSink,
) -> anyhow::Result<(String, u64)> {
if cfg.provider == "broker" {
let sock = cfg
.broker_sock
.as_deref()
.ok_or_else(|| anyhow::anyhow!("broker mode with no socket"))?;
return crate::broker::chat_via_broker(sock, cfg, messages).map(|t| (t, 0));
}
let mut candidates: Vec<AgentConfig> = Vec::new();
for model in crate::tiers::models_for(cfg.provider, task, &cfg.model) {
candidates.push(AgentConfig { model, ..cfg.clone() });
}
for alt in rotation_candidates(cfg.provider) {
for model in crate::tiers::models_for(alt.provider, task, &alt.model) {
candidates.push(AgentConfig { model, ..alt.clone() });
}
}
let mut last: Option<anyhow::Error> = None;
for c in candidates {
match attempt(&c, &messages, task, retrieval, reborrow(&mut sink)) {
Ok(ok) => return Ok(ok),
Err(e) => {
let recoverable = e.downcast_ref::<RateLimited>().is_some()
|| e.downcast_ref::<ModelUnavailable>().is_some();
last = Some(e);
if !recoverable {
break;
}
}
}
}
Err(last.unwrap_or_else(|| anyhow::anyhow!("no model candidates for task '{task}'")))
}
fn attempt(
cfg: &AgentConfig,
messages: &[serde_json::Value],
task: &str,
retrieval: &str,
sink: DeltaSink,
) -> anyhow::Result<(String, u64)> {
let call_id = crate::llm_log::next_call_id();
let start = std::time::Instant::now();
let result = {
let mut raw = String::new();
let mut emitted = 0usize;
let mut wrapped;
let provider_sink: DeltaSink = match sink {
Some(on_delta) => {
wrapped = move |d: &str| {
raw.push_str(d);
let vis = stream_visible(&raw);
if vis.len() > emitted {
on_delta(&vis[emitted..]);
emitted = vis.len();
}
};
Some(&mut wrapped as &mut dyn FnMut(&str))
}
None => None,
};
match cfg.provider {
"anthropic" => chat_anthropic(cfg, messages, provider_sink),
"bedrock" => chat_bedrock(cfg, messages, provider_sink),
_ => chat_openai(cfg, messages, provider_sink),
}
};
let latency_ms = start.elapsed().as_millis() as u64;
match result {
Ok((text, pt, ct)) => {
crate::llm_log::record(&crate::llm_log::CallRecord {
call_id, task, provider: cfg.provider, model: &cfg.model, retrieval,
prompt_tokens: pt, completion_tokens: ct, latency_ms,
ok: true, error: None, input: messages, output: &text,
});
Ok((strip_reasoning(&text), call_id))
}
Err(e) => {
let msg = e.to_string();
crate::llm_log::record(&crate::llm_log::CallRecord {
call_id, task, provider: cfg.provider, model: &cfg.model, retrieval,
prompt_tokens: 0, completion_tokens: 0, latency_ms,
ok: false, error: Some(&msg), input: messages, output: "",
});
Err(e)
}
}
}
pub fn chat(cfg: &AgentConfig, messages: Vec<serde_json::Value>, task: &str) -> anyhow::Result<String> {
chat_with_id(cfg, messages, task, "n/a").map(|(text, _)| text)
}
fn chat_openai(
cfg: &AgentConfig,
messages: &[serde_json::Value],
sink: DeltaSink,
) -> anyhow::Result<(String, u64, u64)> {
let azure = cfg.provider == "azure";
let url = if azure {
cfg.url.clone()
} else {
format!("{}/chat/completions", cfg.url)
};
let mut body = serde_json::json!({
"model": cfg.model,
"messages": messages,
"max_tokens": cfg.max_tokens,
"temperature": cfg.temperature
});
if sink.is_some() {
body["stream"] = serde_json::json!(true);
if matches!(cfg.provider, "openai" | "groq" | "azure") {
body["stream_options"] = serde_json::json!({ "include_usage": true });
}
}
let mut req = ureq::post(&url)
.timeout(std::time::Duration::from_secs(30))
.set("Content-Type", "application/json");
req = if azure {
req.set("api-key", &cfg.key)
} else {
req.set("Authorization", &format!("Bearer {}", cfg.key))
};
let resp = match req.send_json(body) {
Ok(r) => r,
Err(ureq::Error::Status(code, r)) => {
let body = r.into_string().unwrap_or_default();
let api_msg = serde_json::from_str::<serde_json::Value>(&body).ok().and_then(|j| {
let node = if j.is_array() { j[0].clone() } else { j };
node["error"]["message"].as_str().map(str::to_string)
});
let retired = is_retired_model(code, api_msg.as_deref());
let msg = match code {
429 => match api_msg.as_deref().and_then(retry_secs) {
Some(s) => format!(
"rate limit reached — wait ~{s}s and retry (free tier). \
Tip: raise limits, switch model with MARS_LLM_MODEL, or use \
GROQ_API_KEY / a local Ollama via MARS_LLM_URL."
),
None => "rate limit reached (free tier) — wait ~30s and retry, or \
switch model/provider (MARS_LLM_MODEL / GROQ_API_KEY / \
MARS_LLM_URL for local Ollama)."
.to_string(),
},
401 | 403 => format!(
"auth failed — check your API key. ({})",
api_msg.as_deref().unwrap_or("invalid credentials")
),
_ => api_msg.unwrap_or_else(|| format!("HTTP {code}")),
};
if code == 429 {
return Err(anyhow::Error::new(RateLimited(msg)));
}
if retired {
return Err(anyhow::Error::new(ModelUnavailable(msg)));
}
anyhow::bail!("{msg}");
}
Err(e) => anyhow::bail!("{e}"),
};
if let Some(on_delta) = sink {
use std::io::BufRead;
let mut text = String::new();
let (mut pt, mut ct) = (0u64, 0u64);
for line in std::io::BufReader::new(resp.into_reader()).lines() {
let line = line?;
let Some(data) = line.strip_prefix("data:") else { continue };
let data = data.trim();
if data == "[DONE]" {
break;
}
let Ok(j) = serde_json::from_str::<serde_json::Value>(data) else { continue };
if let Some(msg) = j["error"]["message"].as_str() {
anyhow::bail!("{msg}");
}
if let Some(d) = j["choices"][0]["delta"]["content"].as_str() {
if !d.is_empty() {
text.push_str(d);
on_delta(d);
}
}
if j["usage"].is_object() {
pt = j["usage"]["prompt_tokens"].as_u64().unwrap_or(pt);
ct = j["usage"]["completion_tokens"].as_u64().unwrap_or(ct);
}
}
return Ok((text, pt, ct));
}
let json: serde_json::Value = resp.into_json()?;
if let Some(msg) = json["error"]["message"].as_str() {
anyhow::bail!("{msg}");
}
let text = json["choices"][0]["message"]["content"].as_str().unwrap_or("").to_string();
let pt = json["usage"]["prompt_tokens"].as_u64().unwrap_or(0);
let ct = json["usage"]["completion_tokens"].as_u64().unwrap_or(0);
Ok((text, pt, ct))
}
fn chat_anthropic(
cfg: &AgentConfig,
messages: &[serde_json::Value],
sink: DeltaSink,
) -> anyhow::Result<(String, u64, u64)> {
let mut system = String::new();
let mut msgs: Vec<serde_json::Value> = Vec::new();
for m in messages {
if m["role"].as_str() == Some("system") {
if !system.is_empty() {
system.push('\n');
}
system.push_str(m["content"].as_str().unwrap_or(""));
} else {
msgs.push(m.clone());
}
}
let url = format!("{}/v1/messages", cfg.url);
let mut body = serde_json::json!({
"model": cfg.model,
"max_tokens": cfg.max_tokens,
"system": system,
"messages": msgs
});
if sink.is_some() {
body["stream"] = serde_json::json!(true);
}
let resp = match ureq::post(&url)
.timeout(std::time::Duration::from_secs(30))
.set("x-api-key", &cfg.key)
.set("anthropic-version", "2023-06-01")
.set("Content-Type", "application/json")
.send_json(body)
{
Ok(r) => r,
Err(ureq::Error::Status(code, r)) => {
let body = r.into_string().unwrap_or_default();
let api_msg = serde_json::from_str::<serde_json::Value>(&body)
.ok()
.and_then(|j| j["error"]["message"].as_str().map(str::to_string));
let retired = is_retired_model(code, api_msg.as_deref());
let msg = match code {
429 => "rate limit reached (Anthropic) — wait and retry, or switch model \
with MARS_LLM_MODEL."
.to_string(),
401 | 403 => format!(
"auth failed — check ANTHROPIC_API_KEY. ({})",
api_msg.as_deref().unwrap_or("invalid credentials")
),
_ => api_msg.unwrap_or_else(|| format!("HTTP {code}")),
};
if code == 429 {
return Err(anyhow::Error::new(RateLimited(msg)));
}
if retired {
return Err(anyhow::Error::new(ModelUnavailable(msg)));
}
anyhow::bail!("{msg}");
}
Err(e) => anyhow::bail!("{e}"),
};
if let Some(on_delta) = sink {
use std::io::BufRead;
let mut text = String::new();
let (mut pt, mut ct) = (0u64, 0u64);
for line in std::io::BufReader::new(resp.into_reader()).lines() {
let line = line?;
let Some(data) = line.strip_prefix("data:") else { continue };
let Ok(j) = serde_json::from_str::<serde_json::Value>(data.trim()) else { continue };
match j["type"].as_str().unwrap_or("") {
"content_block_delta" => {
if let Some(d) = j["delta"]["text"].as_str() {
if !d.is_empty() {
text.push_str(d);
on_delta(d);
}
}
}
"message_start" => {
pt = j["message"]["usage"]["input_tokens"].as_u64().unwrap_or(0);
}
"message_delta" => {
ct = j["usage"]["output_tokens"].as_u64().unwrap_or(ct);
}
"error" => {
anyhow::bail!("{}", j["error"]["message"].as_str().unwrap_or("stream error"));
}
_ => {}
}
}
return Ok((text, pt, ct));
}
let json: serde_json::Value = resp.into_json()?;
if let Some(msg) = json["error"]["message"].as_str() {
anyhow::bail!("{msg}");
}
let text = json["content"]
.as_array()
.map(|blocks| blocks.iter().filter_map(|b| b["text"].as_str()).collect::<Vec<_>>().join(""))
.unwrap_or_default();
let pt = json["usage"]["input_tokens"].as_u64().unwrap_or(0);
let ct = json["usage"]["output_tokens"].as_u64().unwrap_or(0);
Ok((text, pt, ct))
}
pub fn build_bedrock_body(
messages: &[serde_json::Value],
max_tokens: u32,
temperature: f64,
) -> serde_json::Value {
let mut system: Vec<serde_json::Value> = Vec::new();
let mut msgs: Vec<serde_json::Value> = Vec::new();
for m in messages {
let content = m["content"].as_str().unwrap_or("");
if m["role"].as_str() == Some("system") {
system.push(serde_json::json!({ "text": content }));
} else {
msgs.push(serde_json::json!({
"role": m["role"].as_str().unwrap_or("user"),
"content": [{ "text": content }],
}));
}
}
serde_json::json!({
"system": system,
"messages": msgs,
"inferenceConfig": { "maxTokens": max_tokens, "temperature": temperature },
})
}
fn chat_bedrock(
cfg: &AgentConfig,
messages: &[serde_json::Value],
_sink: DeltaSink,
) -> anyhow::Result<(String, u64, u64)> {
let url = format!("{}/model/{}/converse", cfg.url, cfg.model);
let body = build_bedrock_body(messages, cfg.max_tokens, cfg.temperature);
let resp = match ureq::post(&url)
.timeout(std::time::Duration::from_secs(30))
.set("Authorization", &format!("Bearer {}", cfg.key))
.set("Content-Type", "application/json")
.send_json(body)
{
Ok(r) => r,
Err(ureq::Error::Status(code, r)) => {
let body = r.into_string().unwrap_or_default();
let api_msg = serde_json::from_str::<serde_json::Value>(&body)
.ok()
.and_then(|j| j["message"].as_str().map(str::to_string));
let retired = is_retired_model(code, api_msg.as_deref());
let msg = match code {
429 => "rate limit / throttled (Bedrock) — wait and retry, or switch model \
with MARS_LLM_MODEL."
.to_string(),
401 | 403 => format!(
"auth failed — check AWS_BEARER_TOKEN_BEDROCK and the region/model access. ({})",
api_msg.as_deref().unwrap_or("invalid credentials")
),
_ => api_msg.unwrap_or_else(|| format!("HTTP {code}")),
};
if code == 429 {
return Err(anyhow::Error::new(RateLimited(msg)));
}
if retired {
return Err(anyhow::Error::new(ModelUnavailable(msg)));
}
anyhow::bail!("{msg}");
}
Err(e) => anyhow::bail!("{e}"),
};
let json: serde_json::Value = resp.into_json()?;
if let Some(msg) = json["message"].as_str() {
if json["output"].is_null() {
anyhow::bail!("{msg}");
}
}
let text = json["output"]["message"]["content"]
.as_array()
.map(|blocks| blocks.iter().filter_map(|b| b["text"].as_str()).collect::<Vec<_>>().join(""))
.unwrap_or_default();
let pt = json["usage"]["inputTokens"].as_u64().unwrap_or(0);
let ct = json["usage"]["outputTokens"].as_u64().unwrap_or(0);
Ok((text, pt, ct))
}