mirror of
https://github.com/tinyhumansai/openhuman.git
synced 2026-07-27 21:08:00 +00:00
Merge remote-tracking branch 'upstream/main'
This commit is contained in:
@@ -91,13 +91,8 @@ const handleDeepLinkUrls = async (urls: string[] | null | undefined) => {
|
||||
|
||||
try {
|
||||
const parsed = new URL(url);
|
||||
<<<<<<< HEAD
|
||||
if (parsed.protocol !== 'openhuman:') return;
|
||||
if (parsed.hostname !== 'auth') return;
|
||||
=======
|
||||
if (parsed.protocol !== "outsourced:") return;
|
||||
if (parsed.hostname !== "auth") return;
|
||||
>>>>>>> fix/telegram-mcp
|
||||
|
||||
const token = parsed.searchParams.get("token");
|
||||
if (!token) return;
|
||||
|
||||
@@ -405,10 +405,10 @@ impl Agent {
|
||||
.tools(tools)
|
||||
.memory(memory)
|
||||
.tool_dispatcher(tool_dispatcher)
|
||||
.memory_loader(Box::new(DefaultMemoryLoader::new(
|
||||
5,
|
||||
config.memory.min_relevance_score,
|
||||
)))
|
||||
.memory_loader(Box::new(
|
||||
DefaultMemoryLoader::new(5, config.memory.min_relevance_score)
|
||||
.with_max_chars(config.agent.max_memory_context_chars),
|
||||
))
|
||||
.prompt_builder(prompt_builder)
|
||||
.config(config.agent.clone())
|
||||
.model_name(model_name)
|
||||
@@ -709,6 +709,7 @@ impl Agent {
|
||||
} else {
|
||||
None
|
||||
},
|
||||
system_prompt_cache_boundary: None,
|
||||
},
|
||||
&effective_model,
|
||||
self.temperature,
|
||||
@@ -951,6 +952,7 @@ mod tests {
|
||||
return Ok(crate::openhuman::providers::ChatResponse {
|
||||
text: Some("done".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
});
|
||||
}
|
||||
Ok(guard.remove(0))
|
||||
@@ -994,6 +996,7 @@ mod tests {
|
||||
responses: Mutex::new(vec![crate::openhuman::providers::ChatResponse {
|
||||
text: Some("hello".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
}]),
|
||||
});
|
||||
|
||||
@@ -1032,10 +1035,12 @@ mod tests {
|
||||
name: "echo".into(),
|
||||
arguments: "{}".into(),
|
||||
}],
|
||||
usage: None,
|
||||
},
|
||||
crate::openhuman::providers::ChatResponse {
|
||||
text: Some("done".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
},
|
||||
]),
|
||||
});
|
||||
@@ -1078,10 +1083,12 @@ mod tests {
|
||||
.into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
},
|
||||
crate::openhuman::providers::ChatResponse {
|
||||
text: Some("done".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
},
|
||||
]),
|
||||
});
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
//! Token cost tracking for agent loop budget enforcement.
|
||||
//!
|
||||
//! Tracks cumulative token usage across inference calls and enforces
|
||||
//! the `max_cost_per_day_cents` budget from the security policy.
|
||||
|
||||
use crate::openhuman::providers::UsageInfo;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
/// Per-token pricing in microdollars (millionths of a dollar).
|
||||
/// Default pricing is conservative; callers should provide model-specific rates.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TokenPricing {
|
||||
/// Cost per input token in microdollars.
|
||||
pub input_token_microdollars: u64,
|
||||
/// Cost per output token in microdollars.
|
||||
pub output_token_microdollars: u64,
|
||||
}
|
||||
|
||||
impl Default for TokenPricing {
|
||||
fn default() -> Self {
|
||||
// Conservative defaults (~$3/1M input, ~$15/1M output — Sonnet-class)
|
||||
Self {
|
||||
input_token_microdollars: 3,
|
||||
output_token_microdollars: 15,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Thread-safe cumulative token and cost tracker.
|
||||
#[derive(Debug)]
|
||||
pub struct CostTracker {
|
||||
total_input_tokens: AtomicU64,
|
||||
total_output_tokens: AtomicU64,
|
||||
total_cost_microdollars: AtomicU64,
|
||||
pricing: TokenPricing,
|
||||
/// Budget in microdollars (0 = unlimited).
|
||||
budget_microdollars: u64,
|
||||
}
|
||||
|
||||
impl CostTracker {
|
||||
/// Create a new tracker with the given pricing and budget (in cents).
|
||||
pub fn new(pricing: TokenPricing, budget_cents: u32) -> Self {
|
||||
Self {
|
||||
total_input_tokens: AtomicU64::new(0),
|
||||
total_output_tokens: AtomicU64::new(0),
|
||||
total_cost_microdollars: AtomicU64::new(0),
|
||||
pricing,
|
||||
budget_microdollars: budget_cents as u64 * 10_000, // cents → microdollars
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a tracker with default pricing and a budget in cents.
|
||||
pub fn with_budget_cents(budget_cents: u32) -> Self {
|
||||
Self::new(TokenPricing::default(), budget_cents)
|
||||
}
|
||||
|
||||
/// Record usage from a provider response.
|
||||
pub fn record_usage(&self, usage: &UsageInfo) {
|
||||
self.total_input_tokens
|
||||
.fetch_add(usage.input_tokens, Ordering::Relaxed);
|
||||
self.total_output_tokens
|
||||
.fetch_add(usage.output_tokens, Ordering::Relaxed);
|
||||
|
||||
let cost = usage.input_tokens * self.pricing.input_token_microdollars
|
||||
+ usage.output_tokens * self.pricing.output_token_microdollars;
|
||||
self.total_cost_microdollars
|
||||
.fetch_add(cost, Ordering::Relaxed);
|
||||
|
||||
tracing::debug!(
|
||||
input_tokens = usage.input_tokens,
|
||||
output_tokens = usage.output_tokens,
|
||||
cost_microdollars = cost,
|
||||
total_cost_microdollars = self.total_cost_microdollars.load(Ordering::Relaxed),
|
||||
"[cost_tracker] recorded usage"
|
||||
);
|
||||
}
|
||||
|
||||
/// Check whether the budget has been exceeded.
|
||||
/// Returns `Ok(())` if within budget, or `Err` with spent/budget amounts.
|
||||
pub fn check_budget(&self) -> Result<(), (u64, u64)> {
|
||||
if self.budget_microdollars == 0 {
|
||||
return Ok(()); // Unlimited
|
||||
}
|
||||
let spent = self.total_cost_microdollars.load(Ordering::Relaxed);
|
||||
if spent > self.budget_microdollars {
|
||||
Err((spent, self.budget_microdollars))
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the total input tokens recorded.
|
||||
pub fn total_input_tokens(&self) -> u64 {
|
||||
self.total_input_tokens.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
/// Get the total output tokens recorded.
|
||||
pub fn total_output_tokens(&self) -> u64 {
|
||||
self.total_output_tokens.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
/// Get the total cost in microdollars.
|
||||
pub fn total_cost_microdollars(&self) -> u64 {
|
||||
self.total_cost_microdollars.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
/// Human-readable cost summary.
|
||||
pub fn summary(&self) -> String {
|
||||
let input = self.total_input_tokens.load(Ordering::Relaxed);
|
||||
let output = self.total_output_tokens.load(Ordering::Relaxed);
|
||||
let cost = self.total_cost_microdollars.load(Ordering::Relaxed);
|
||||
let dollars = cost as f64 / 1_000_000.0;
|
||||
format!("Tokens: {input} in / {output} out | Cost: ${dollars:.4}")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn tracks_cumulative_usage() {
|
||||
let tracker = CostTracker::with_budget_cents(1000);
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
context_window: 0,
|
||||
});
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 2000,
|
||||
output_tokens: 1000,
|
||||
context_window: 0,
|
||||
});
|
||||
|
||||
assert_eq!(tracker.total_input_tokens(), 3000);
|
||||
assert_eq!(tracker.total_output_tokens(), 1500);
|
||||
assert!(tracker.total_cost_microdollars() > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn budget_enforcement() {
|
||||
// Budget: 1 cent = 10,000 microdollars
|
||||
let tracker = CostTracker::new(
|
||||
TokenPricing {
|
||||
input_token_microdollars: 100,
|
||||
output_token_microdollars: 100,
|
||||
},
|
||||
1, // 1 cent
|
||||
);
|
||||
|
||||
// 50 input + 50 output = 100 tokens × 100 = 10,000 microdollars = 1 cent (at limit)
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 50,
|
||||
output_tokens: 50,
|
||||
context_window: 0,
|
||||
});
|
||||
assert!(tracker.check_budget().is_ok());
|
||||
|
||||
// One more token pushes over budget
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 1,
|
||||
output_tokens: 0,
|
||||
context_window: 0,
|
||||
});
|
||||
assert!(tracker.check_budget().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unlimited_budget() {
|
||||
let tracker = CostTracker::with_budget_cents(0);
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 1_000_000,
|
||||
output_tokens: 1_000_000,
|
||||
context_window: 0,
|
||||
});
|
||||
assert!(tracker.check_budget().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn summary_format() {
|
||||
let tracker = CostTracker::with_budget_cents(100);
|
||||
tracker.record_usage(&UsageInfo {
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
context_window: 0,
|
||||
});
|
||||
let summary = tracker.summary();
|
||||
assert!(summary.contains("1000 in"));
|
||||
assert!(summary.contains("500 out"));
|
||||
assert!(summary.contains("$"));
|
||||
}
|
||||
}
|
||||
@@ -255,6 +255,7 @@ mod tests {
|
||||
.into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
let dispatcher = XmlToolDispatcher;
|
||||
let (_, calls) = dispatcher.parse_response(&response);
|
||||
@@ -271,6 +272,7 @@ mod tests {
|
||||
name: "file_read".into(),
|
||||
arguments: "{\"path\":\"a.txt\"}".into(),
|
||||
}],
|
||||
usage: None,
|
||||
};
|
||||
let dispatcher = NativeToolDispatcher;
|
||||
let (_, calls) = dispatcher.parse_response(&response);
|
||||
@@ -300,6 +302,7 @@ mod tests {
|
||||
.into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
let dispatcher = NativeToolDispatcher;
|
||||
let (text, calls) = dispatcher.parse_response(&response);
|
||||
@@ -316,6 +319,7 @@ mod tests {
|
||||
"Let me run this.\n<invoke>{\"name\":\"shell\",\"arguments\":{\"command\":\"pwd\"}}</invoke>".into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
let dispatcher = NativeToolDispatcher;
|
||||
let (text, calls) = dispatcher.parse_response(&response);
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
//! Structured error types for the agent loop.
|
||||
//!
|
||||
//! Replaces generic `anyhow::bail!` with typed variants so callers can
|
||||
//! distinguish retryable errors from permanent failures and take appropriate
|
||||
//! recovery actions (e.g. triggering compaction on context-limit errors).
|
||||
|
||||
use std::fmt;
|
||||
|
||||
/// Structured error type for agent loop operations.
|
||||
#[derive(Debug)]
|
||||
pub enum AgentError {
|
||||
/// The LLM provider returned an error.
|
||||
ProviderError { message: String, retryable: bool },
|
||||
/// Context window is exhausted and compaction cannot help.
|
||||
ContextLimitExceeded { utilization_pct: u8 },
|
||||
/// A tool execution failed.
|
||||
ToolExecutionError { tool_name: String, message: String },
|
||||
/// The daily cost budget has been exceeded.
|
||||
CostBudgetExceeded {
|
||||
spent_microdollars: u64,
|
||||
budget_microdollars: u64,
|
||||
},
|
||||
/// The agent exceeded its maximum tool iterations.
|
||||
MaxIterationsExceeded { max: usize },
|
||||
/// History compaction failed.
|
||||
CompactionFailed {
|
||||
message: String,
|
||||
consecutive_failures: u8,
|
||||
},
|
||||
/// Channel permission denied for a tool operation.
|
||||
PermissionDenied {
|
||||
tool_name: String,
|
||||
required_level: String,
|
||||
channel_max_level: String,
|
||||
},
|
||||
/// Generic/untyped error (escape hatch for migration).
|
||||
Other(anyhow::Error),
|
||||
}
|
||||
|
||||
impl fmt::Display for AgentError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::ProviderError { message, retryable } => {
|
||||
write!(f, "Provider error (retryable={retryable}): {message}")
|
||||
}
|
||||
Self::ContextLimitExceeded { utilization_pct } => {
|
||||
write!(
|
||||
f,
|
||||
"Context window exhausted ({utilization_pct}% utilized, compaction disabled)"
|
||||
)
|
||||
}
|
||||
Self::ToolExecutionError { tool_name, message } => {
|
||||
write!(f, "Tool execution error [{tool_name}]: {message}")
|
||||
}
|
||||
Self::CostBudgetExceeded {
|
||||
spent_microdollars,
|
||||
budget_microdollars,
|
||||
} => {
|
||||
let spent = *spent_microdollars as f64 / 1_000_000.0;
|
||||
let budget = *budget_microdollars as f64 / 1_000_000.0;
|
||||
write!(
|
||||
f,
|
||||
"Daily cost budget exceeded: spent ${spent:.4}, budget ${budget:.4}"
|
||||
)
|
||||
}
|
||||
Self::MaxIterationsExceeded { max } => {
|
||||
write!(f, "Agent exceeded maximum tool iterations ({max})")
|
||||
}
|
||||
Self::CompactionFailed {
|
||||
message,
|
||||
consecutive_failures,
|
||||
} => {
|
||||
write!(
|
||||
f,
|
||||
"Compaction failed ({consecutive_failures} consecutive): {message}"
|
||||
)
|
||||
}
|
||||
Self::PermissionDenied {
|
||||
tool_name,
|
||||
required_level,
|
||||
channel_max_level,
|
||||
} => {
|
||||
write!(
|
||||
f,
|
||||
"Permission denied for tool '{tool_name}': requires {required_level}, channel allows {channel_max_level}"
|
||||
)
|
||||
}
|
||||
Self::Other(e) => write!(f, "{e}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for AgentError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Other(e) => Some(e.as_ref()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<anyhow::Error> for AgentError {
|
||||
fn from(e: anyhow::Error) -> Self {
|
||||
// Attempt to recover a typed AgentError that was wrapped in anyhow.
|
||||
match e.downcast::<AgentError>() {
|
||||
Ok(agent_err) => agent_err,
|
||||
Err(other) => Self::Other(other),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if an error message indicates a context/prompt-too-long failure.
|
||||
pub fn is_context_limit_error(error_msg: &str) -> bool {
|
||||
let lower = error_msg.to_lowercase();
|
||||
lower.contains("prompt is too long")
|
||||
|| lower.contains("context_length_exceeded")
|
||||
|| lower.contains("maximum context length")
|
||||
|| lower.contains("prompt too long")
|
||||
|| lower.contains("token limit")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn display_formatting() {
|
||||
let err = AgentError::MaxIterationsExceeded { max: 10 };
|
||||
assert_eq!(
|
||||
err.to_string(),
|
||||
"Agent exceeded maximum tool iterations (10)"
|
||||
);
|
||||
|
||||
let err = AgentError::CostBudgetExceeded {
|
||||
spent_microdollars: 5_500_000,
|
||||
budget_microdollars: 5_000_000,
|
||||
};
|
||||
assert!(err.to_string().contains("5.5000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_limit_detection() {
|
||||
assert!(is_context_limit_error("prompt is too long for model"));
|
||||
assert!(is_context_limit_error("context_length_exceeded"));
|
||||
assert!(!is_context_limit_error("rate limit exceeded"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn permission_denied_display() {
|
||||
let err = AgentError::PermissionDenied {
|
||||
tool_name: "shell".into(),
|
||||
required_level: "Execute".into(),
|
||||
channel_max_level: "ReadOnly".into(),
|
||||
};
|
||||
assert!(err.to_string().contains("shell"));
|
||||
assert!(err.to_string().contains("Execute"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,216 @@
|
||||
//! Typed event system for agent loop observability.
|
||||
//!
|
||||
//! Replaces the basic `ToolEventObserver` with a comprehensive `AgentEvent`
|
||||
//! enum broadcast via `tokio::sync::broadcast`. Multiple consumers (Socket.IO
|
||||
//! relay, logging, cost tracking) can subscribe to the same event stream.
|
||||
|
||||
use crate::openhuman::providers::UsageInfo;
|
||||
|
||||
/// Events emitted during agent loop execution.
|
||||
///
|
||||
/// Subscribers receive these via `tokio::sync::broadcast::Receiver<AgentEvent>`.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum AgentEvent {
|
||||
/// An LLM inference call is about to be made.
|
||||
InferenceStart {
|
||||
iteration: usize,
|
||||
message_count: usize,
|
||||
},
|
||||
|
||||
/// An LLM inference call completed.
|
||||
InferenceComplete {
|
||||
iteration: usize,
|
||||
has_tool_calls: bool,
|
||||
usage: Option<UsageInfo>,
|
||||
},
|
||||
|
||||
/// Tool calls were parsed from the LLM response.
|
||||
ToolCallsParsed {
|
||||
tool_names: Vec<String>,
|
||||
/// Full arguments per tool call (parallel with tool_names).
|
||||
tool_arguments: Vec<serde_json::Value>,
|
||||
/// Optional tool_call_id per call (parallel with tool_names).
|
||||
tool_call_ids: Vec<Option<String>>,
|
||||
iteration: usize,
|
||||
},
|
||||
|
||||
/// A single tool execution is starting.
|
||||
ToolExecutionStart { name: String, iteration: usize },
|
||||
|
||||
/// A single tool execution completed.
|
||||
ToolExecutionComplete {
|
||||
name: String,
|
||||
/// The actual tool output string.
|
||||
output: String,
|
||||
output_chars: usize,
|
||||
elapsed_ms: u64,
|
||||
success: bool,
|
||||
tool_call_id: Option<String>,
|
||||
iteration: usize,
|
||||
},
|
||||
|
||||
/// Context compaction was triggered.
|
||||
CompactionTriggered {
|
||||
messages_before: usize,
|
||||
messages_after: usize,
|
||||
},
|
||||
|
||||
/// Context compaction failed.
|
||||
CompactionFailed {
|
||||
error: String,
|
||||
consecutive_failures: u8,
|
||||
},
|
||||
|
||||
/// The agent turn completed with a final text response.
|
||||
TurnComplete {
|
||||
text_chars: usize,
|
||||
total_iterations: usize,
|
||||
},
|
||||
|
||||
/// An error occurred during the agent loop.
|
||||
Error { message: String, recoverable: bool },
|
||||
|
||||
/// Cost update after an inference call.
|
||||
CostUpdate {
|
||||
total_input_tokens: u64,
|
||||
total_output_tokens: u64,
|
||||
total_cost_microdollars: u64,
|
||||
},
|
||||
}
|
||||
|
||||
/// Convenience sender wrapper that silently drops events if no receivers are listening.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct EventSender {
|
||||
tx: tokio::sync::broadcast::Sender<AgentEvent>,
|
||||
}
|
||||
|
||||
impl EventSender {
|
||||
/// Create a new event sender with the given channel capacity.
|
||||
/// Capacity is clamped to at least 1 to avoid a broadcast channel panic.
|
||||
pub fn new(capacity: usize) -> (Self, tokio::sync::broadcast::Receiver<AgentEvent>) {
|
||||
let cap = capacity.max(1);
|
||||
let (tx, rx) = tokio::sync::broadcast::channel(cap);
|
||||
(Self { tx }, rx)
|
||||
}
|
||||
|
||||
/// Emit an event. Silently drops if no receivers are listening.
|
||||
pub fn emit(&self, event: AgentEvent) {
|
||||
tracing::trace!(
|
||||
event = ?std::mem::discriminant(&event),
|
||||
receivers = self.tx.receiver_count(),
|
||||
"[agent_events] emitting event"
|
||||
);
|
||||
let _ = self.tx.send(event);
|
||||
}
|
||||
|
||||
/// Subscribe to the event stream.
|
||||
pub fn subscribe(&self) -> tokio::sync::broadcast::Receiver<AgentEvent> {
|
||||
self.tx.subscribe()
|
||||
}
|
||||
}
|
||||
|
||||
/// Default broadcast channel capacity for agent events.
|
||||
pub const DEFAULT_EVENT_CHANNEL_CAPACITY: usize = 128;
|
||||
|
||||
/// Bridge adapter that converts `AgentEvent`s into `ToolEventObserver` callbacks,
|
||||
/// allowing gradual migration from the old observer pattern.
|
||||
pub struct ObserverBridge {
|
||||
observer: std::sync::Arc<dyn super::observer::ToolEventObserver>,
|
||||
}
|
||||
|
||||
impl ObserverBridge {
|
||||
pub fn new(observer: std::sync::Arc<dyn super::observer::ToolEventObserver>) -> Self {
|
||||
Self { observer }
|
||||
}
|
||||
|
||||
/// Process an event and forward to the legacy observer if applicable.
|
||||
pub fn handle_event(&self, event: &AgentEvent) {
|
||||
tracing::trace!(
|
||||
event = ?std::mem::discriminant(event),
|
||||
"[agent_events] ObserverBridge handling event"
|
||||
);
|
||||
match event {
|
||||
AgentEvent::ToolCallsParsed {
|
||||
tool_names,
|
||||
tool_arguments,
|
||||
tool_call_ids,
|
||||
iteration,
|
||||
} => {
|
||||
let calls: Vec<super::dispatcher::ParsedToolCall> = tool_names
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, name)| super::dispatcher::ParsedToolCall {
|
||||
name: name.clone(),
|
||||
arguments: tool_arguments
|
||||
.get(i)
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
tool_call_id: tool_call_ids.get(i).cloned().flatten(),
|
||||
})
|
||||
.collect();
|
||||
self.observer.on_tool_calls(&calls, *iteration as u32);
|
||||
}
|
||||
AgentEvent::ToolExecutionComplete {
|
||||
name,
|
||||
output,
|
||||
success,
|
||||
tool_call_id,
|
||||
iteration,
|
||||
..
|
||||
} => {
|
||||
let results = vec![super::dispatcher::ToolExecutionResult {
|
||||
name: name.clone(),
|
||||
output: output.clone(),
|
||||
success: *success,
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
}];
|
||||
self.observer.on_tool_results(&results, *iteration as u32);
|
||||
}
|
||||
_ => {} // Other events have no legacy equivalent
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn event_sender_works_without_receivers() {
|
||||
let (sender, _rx) = EventSender::new(16);
|
||||
// Should not panic even with no active receivers
|
||||
drop(_rx);
|
||||
sender.emit(AgentEvent::TurnComplete {
|
||||
text_chars: 100,
|
||||
total_iterations: 1,
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_sender_delivers_to_subscriber() {
|
||||
let (sender, mut rx) = EventSender::new(16);
|
||||
sender.emit(AgentEvent::InferenceStart {
|
||||
iteration: 1,
|
||||
message_count: 5,
|
||||
});
|
||||
let event = rx.try_recv().unwrap();
|
||||
assert!(matches!(
|
||||
event,
|
||||
AgentEvent::InferenceStart { iteration: 1, .. }
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn multiple_subscribers_receive_events() {
|
||||
let (sender, mut rx1) = EventSender::new(16);
|
||||
let mut rx2 = sender.subscribe();
|
||||
|
||||
sender.emit(AgentEvent::TurnComplete {
|
||||
text_chars: 42,
|
||||
total_iterations: 2,
|
||||
});
|
||||
|
||||
assert!(rx1.try_recv().is_ok());
|
||||
assert!(rx2.try_recv().is_ok());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,217 @@
|
||||
//! Pre-inference context window guard with compaction circuit breaker.
|
||||
//!
|
||||
//! Checks context utilization before each LLM call and triggers auto-compaction
|
||||
//! when usage exceeds a threshold. A circuit breaker disables compaction after
|
||||
//! consecutive failures to prevent infinite retry loops.
|
||||
|
||||
use crate::openhuman::providers::UsageInfo;
|
||||
|
||||
/// Threshold (0.0–1.0) at which auto-compaction is triggered.
|
||||
const COMPACTION_TRIGGER_THRESHOLD: f64 = 0.90;
|
||||
|
||||
/// Threshold above which, if compaction is disabled, the guard returns an error.
|
||||
const HARD_LIMIT_THRESHOLD: f64 = 0.95;
|
||||
|
||||
/// Number of consecutive compaction failures before the circuit breaker trips.
|
||||
const MAX_CONSECUTIVE_FAILURES: u8 = 3;
|
||||
|
||||
/// Outcome of a pre-inference context check.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum ContextCheckResult {
|
||||
/// Context utilization is within safe limits.
|
||||
Ok,
|
||||
/// Context is near capacity; compaction should be attempted.
|
||||
CompactionNeeded,
|
||||
/// Context is critically full and compaction is disabled (circuit breaker tripped).
|
||||
ContextExhausted { utilization_pct: u8, reason: String },
|
||||
}
|
||||
|
||||
/// Tracks context window utilization and compaction health.
|
||||
#[derive(Debug)]
|
||||
pub struct ContextGuard {
|
||||
/// Last known input token count from the provider.
|
||||
last_input_tokens: u64,
|
||||
/// Last known output token count from the provider.
|
||||
last_output_tokens: u64,
|
||||
/// Model context window size (0 = unknown, guard is a no-op).
|
||||
context_window: u64,
|
||||
/// Number of consecutive compaction failures.
|
||||
consecutive_compaction_failures: u8,
|
||||
/// Whether compaction has been disabled by the circuit breaker.
|
||||
compaction_disabled: bool,
|
||||
}
|
||||
|
||||
impl Default for ContextGuard {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl ContextGuard {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
last_input_tokens: 0,
|
||||
last_output_tokens: 0,
|
||||
context_window: 0,
|
||||
consecutive_compaction_failures: 0,
|
||||
compaction_disabled: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a guard with a known context window size.
|
||||
pub fn with_context_window(context_window: u64) -> Self {
|
||||
Self {
|
||||
context_window,
|
||||
..Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// Update the guard with usage info from the latest provider response.
|
||||
pub fn update_usage(&mut self, usage: &UsageInfo) {
|
||||
self.last_input_tokens = usage.input_tokens;
|
||||
self.last_output_tokens = usage.output_tokens;
|
||||
if usage.context_window > 0 {
|
||||
self.context_window = usage.context_window;
|
||||
}
|
||||
}
|
||||
|
||||
/// Estimate current context utilization as a fraction (0.0–1.0).
|
||||
/// Returns `None` if context window is unknown.
|
||||
pub fn utilization(&self) -> Option<f64> {
|
||||
if self.context_window == 0 {
|
||||
return None;
|
||||
}
|
||||
let total_used = self.last_input_tokens + self.last_output_tokens;
|
||||
Some(total_used as f64 / self.context_window as f64)
|
||||
}
|
||||
|
||||
/// Check whether the context is safe to proceed with another inference call.
|
||||
pub fn check(&self) -> ContextCheckResult {
|
||||
let utilization = match self.utilization() {
|
||||
Some(u) => u,
|
||||
None => return ContextCheckResult::Ok, // Unknown window = no guard
|
||||
};
|
||||
|
||||
if utilization >= HARD_LIMIT_THRESHOLD && self.compaction_disabled {
|
||||
return ContextCheckResult::ContextExhausted {
|
||||
utilization_pct: (utilization * 100.0) as u8,
|
||||
reason: format!(
|
||||
"Context {:.0}% full; compaction disabled after {} consecutive failures",
|
||||
utilization * 100.0,
|
||||
self.consecutive_compaction_failures
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
if utilization >= COMPACTION_TRIGGER_THRESHOLD && !self.compaction_disabled {
|
||||
return ContextCheckResult::CompactionNeeded;
|
||||
}
|
||||
|
||||
ContextCheckResult::Ok
|
||||
}
|
||||
|
||||
/// Record a successful compaction, resetting the failure counter.
|
||||
pub fn record_compaction_success(&mut self) {
|
||||
self.consecutive_compaction_failures = 0;
|
||||
self.compaction_disabled = false;
|
||||
tracing::debug!("[context_guard] compaction succeeded, circuit breaker reset");
|
||||
}
|
||||
|
||||
/// Record a failed compaction attempt. Trips the circuit breaker after
|
||||
/// `MAX_CONSECUTIVE_FAILURES` failures.
|
||||
pub fn record_compaction_failure(&mut self) {
|
||||
self.consecutive_compaction_failures += 1;
|
||||
if self.consecutive_compaction_failures >= MAX_CONSECUTIVE_FAILURES {
|
||||
self.compaction_disabled = true;
|
||||
tracing::warn!(
|
||||
consecutive_failures = self.consecutive_compaction_failures,
|
||||
"[context_guard] circuit breaker tripped — compaction disabled"
|
||||
);
|
||||
} else {
|
||||
tracing::debug!(
|
||||
consecutive_failures = self.consecutive_compaction_failures,
|
||||
max = MAX_CONSECUTIVE_FAILURES,
|
||||
"[context_guard] compaction failed, circuit breaker pending"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Whether the compaction circuit breaker is currently tripped.
|
||||
pub fn is_compaction_disabled(&self) -> bool {
|
||||
self.compaction_disabled
|
||||
}
|
||||
|
||||
/// Number of consecutive compaction failures.
|
||||
pub fn consecutive_failures(&self) -> u8 {
|
||||
self.consecutive_compaction_failures
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn unknown_context_window_always_ok() {
|
||||
let guard = ContextGuard::new();
|
||||
assert_eq!(guard.check(), ContextCheckResult::Ok);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn low_utilization_is_ok() {
|
||||
let mut guard = ContextGuard::with_context_window(100_000);
|
||||
guard.update_usage(&UsageInfo {
|
||||
input_tokens: 10_000,
|
||||
output_tokens: 5_000,
|
||||
context_window: 100_000,
|
||||
});
|
||||
assert_eq!(guard.check(), ContextCheckResult::Ok);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn high_utilization_triggers_compaction() {
|
||||
let mut guard = ContextGuard::with_context_window(100_000);
|
||||
guard.update_usage(&UsageInfo {
|
||||
input_tokens: 85_000,
|
||||
output_tokens: 6_000,
|
||||
context_window: 100_000,
|
||||
});
|
||||
assert_eq!(guard.check(), ContextCheckResult::CompactionNeeded);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn circuit_breaker_trips_after_three_failures() {
|
||||
let mut guard = ContextGuard::with_context_window(100_000);
|
||||
guard.update_usage(&UsageInfo {
|
||||
input_tokens: 90_000,
|
||||
output_tokens: 6_000,
|
||||
context_window: 100_000,
|
||||
});
|
||||
|
||||
guard.record_compaction_failure();
|
||||
guard.record_compaction_failure();
|
||||
assert!(!guard.is_compaction_disabled());
|
||||
|
||||
guard.record_compaction_failure();
|
||||
assert!(guard.is_compaction_disabled());
|
||||
|
||||
// Now at >95%, should return exhausted
|
||||
assert!(matches!(
|
||||
guard.check(),
|
||||
ContextCheckResult::ContextExhausted { .. }
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn success_resets_circuit_breaker() {
|
||||
let mut guard = ContextGuard::with_context_window(100_000);
|
||||
guard.record_compaction_failure();
|
||||
guard.record_compaction_failure();
|
||||
guard.record_compaction_failure();
|
||||
assert!(guard.is_compaction_disabled());
|
||||
|
||||
guard.record_compaction_success();
|
||||
assert!(!guard.is_compaction_disabled());
|
||||
assert_eq!(guard.consecutive_failures(), 0);
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
//! Agent loop: tool-call execution, CLI session, and channel message handling.
|
||||
|
||||
pub(crate) mod context_guard;
|
||||
mod credentials;
|
||||
mod history;
|
||||
mod instructions;
|
||||
|
||||
@@ -100,6 +100,7 @@ impl Provider for VisionProvider {
|
||||
Ok(ChatResponse {
|
||||
text: Some("vision-ok".to_string()),
|
||||
tool_calls: Vec::new(),
|
||||
usage: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -107,6 +107,7 @@ pub(crate) async fn run_tool_call_loop(
|
||||
ChatRequest {
|
||||
messages: &prepared_messages.messages,
|
||||
tools: request_tools,
|
||||
system_prompt_cache_boundary: None,
|
||||
},
|
||||
model,
|
||||
temperature,
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use crate::openhuman::memory::Memory;
|
||||
use async_trait::async_trait;
|
||||
use std::fmt::Write;
|
||||
|
||||
#[async_trait]
|
||||
pub trait MemoryLoader: Send + Sync {
|
||||
@@ -11,6 +10,8 @@ pub trait MemoryLoader: Send + Sync {
|
||||
pub struct DefaultMemoryLoader {
|
||||
limit: usize,
|
||||
min_relevance_score: f64,
|
||||
/// Maximum characters of memory context to inject (0 = unlimited).
|
||||
max_context_chars: usize,
|
||||
}
|
||||
|
||||
impl Default for DefaultMemoryLoader {
|
||||
@@ -18,6 +19,7 @@ impl Default for DefaultMemoryLoader {
|
||||
Self {
|
||||
limit: 5,
|
||||
min_relevance_score: 0.4,
|
||||
max_context_chars: 2000,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -27,8 +29,14 @@ impl DefaultMemoryLoader {
|
||||
Self {
|
||||
limit: limit.max(1),
|
||||
min_relevance_score,
|
||||
max_context_chars: 2000,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_max_chars(mut self, max_chars: usize) -> Self {
|
||||
self.max_context_chars = max_chars;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -43,18 +51,35 @@ impl MemoryLoader for DefaultMemoryLoader {
|
||||
return Ok(String::new());
|
||||
}
|
||||
|
||||
let mut context = String::from("[Memory context]\n");
|
||||
let header = "[Memory context]\n";
|
||||
let mut context = String::from(header);
|
||||
let budget = if self.max_context_chars > 0 {
|
||||
self.max_context_chars
|
||||
} else {
|
||||
usize::MAX
|
||||
};
|
||||
|
||||
for entry in entries {
|
||||
if let Some(score) = entry.score {
|
||||
if score < self.min_relevance_score {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
let _ = writeln!(context, "- {}: {}", entry.key, entry.content);
|
||||
let line = format!("- {}: {}\n", entry.key, entry.content);
|
||||
if context.len() + line.len() > budget {
|
||||
tracing::debug!(
|
||||
budget,
|
||||
current_len = context.len(),
|
||||
skipped_line_len = line.len(),
|
||||
"[memory_loader] context budget reached, skipping remaining entries"
|
||||
);
|
||||
break;
|
||||
}
|
||||
context.push_str(&line);
|
||||
}
|
||||
|
||||
// If all entries were below threshold, return empty
|
||||
if context == "[Memory context]\n" {
|
||||
if context == header {
|
||||
return Ok(String::new());
|
||||
}
|
||||
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
#[allow(clippy::module_inception)]
|
||||
pub mod agent;
|
||||
pub mod classifier;
|
||||
pub mod cost;
|
||||
pub mod dispatcher;
|
||||
pub mod error;
|
||||
pub mod events;
|
||||
pub mod hooks;
|
||||
pub mod host_runtime;
|
||||
pub mod identity;
|
||||
pub mod loop_;
|
||||
pub mod memory_loader;
|
||||
pub mod multimodal;
|
||||
pub mod observer;
|
||||
pub mod prompt;
|
||||
mod schemas;
|
||||
pub mod traits;
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
//! Legacy observer trait for tool events.
|
||||
//!
|
||||
//! Deprecated: prefer the typed `AgentEvent` system in `events.rs`.
|
||||
//! This module is kept for backward compatibility; use `events::ObserverBridge`
|
||||
//! to connect legacy observers to the new event stream.
|
||||
|
||||
use super::dispatcher::{ParsedToolCall, ToolExecutionResult};
|
||||
|
||||
/// Observer for tool events emitted during the agent loop.
|
||||
///
|
||||
/// Implementors receive callbacks as tool calls are parsed and executed,
|
||||
/// enabling real-time event publishing (e.g. to Socket.IO) rather than
|
||||
/// batch-publishing after the entire loop completes.
|
||||
///
|
||||
/// **Deprecated**: Use `AgentEvent` broadcast channel from `events.rs` instead.
|
||||
pub trait ToolEventObserver: Send + Sync {
|
||||
/// Called after tool calls are parsed from the LLM response, before execution.
|
||||
fn on_tool_calls(&self, calls: &[ParsedToolCall], round: u32);
|
||||
|
||||
/// Called after all tool calls in a round have been executed.
|
||||
fn on_tool_results(&self, results: &[ToolExecutionResult], round: u32);
|
||||
}
|
||||
@@ -63,11 +63,19 @@ impl SystemPromptBuilder {
|
||||
|
||||
pub fn build(&self, ctx: &PromptContext<'_>) -> Result<String> {
|
||||
let mut output = String::new();
|
||||
let mut cache_boundary_inserted = false;
|
||||
for section in &self.sections {
|
||||
let part = section.build(ctx)?;
|
||||
if part.trim().is_empty() {
|
||||
continue;
|
||||
}
|
||||
// Insert cache boundary marker before the first dynamic section.
|
||||
// Static sections (identity, tools, safety, skills) are cacheable;
|
||||
// dynamic sections (workspace, datetime, runtime) change per request.
|
||||
if !cache_boundary_inserted && is_dynamic_section(section.name()) {
|
||||
output.push_str("<!-- CACHE_BOUNDARY -->\n\n");
|
||||
cache_boundary_inserted = true;
|
||||
}
|
||||
output.push_str(part.trim_end());
|
||||
output.push_str("\n\n");
|
||||
}
|
||||
@@ -235,6 +243,13 @@ impl PromptSection for DateTimeSection {
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns true for sections whose content changes between requests.
|
||||
/// Static sections (identity, tools, safety, skills) are placed before
|
||||
/// the cache boundary; dynamic sections (workspace, datetime, runtime) after.
|
||||
fn is_dynamic_section(name: &str) -> bool {
|
||||
matches!(name, "workspace" | "datetime" | "runtime")
|
||||
}
|
||||
|
||||
fn inject_workspace_file(prompt: &mut String, workspace_dir: &Path, filename: &str) {
|
||||
let path = workspace_dir.join(filename);
|
||||
if !path.exists() {
|
||||
|
||||
@@ -92,6 +92,7 @@ impl Provider for ScriptedProvider {
|
||||
return Ok(ChatResponse {
|
||||
text: Some("done".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
});
|
||||
}
|
||||
Ok(guard.remove(0))
|
||||
@@ -333,6 +334,7 @@ fn tool_response(calls: Vec<ToolCall>) -> ChatResponse {
|
||||
ChatResponse {
|
||||
text: Some(String::new()),
|
||||
tool_calls: calls,
|
||||
usage: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -341,6 +343,7 @@ fn text_response(text: &str) -> ChatResponse {
|
||||
ChatResponse {
|
||||
text: Some(text.into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -351,6 +354,7 @@ fn xml_tool_response(name: &str, args: &str) -> ChatResponse {
|
||||
"<tool_call>\n{{\"name\": \"{name}\", \"arguments\": {args}}}\n</tool_call>"
|
||||
)),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -725,6 +729,7 @@ async fn turn_handles_empty_text_response() {
|
||||
let provider = Box::new(ScriptedProvider::new(vec![ChatResponse {
|
||||
text: Some(String::new()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
}]));
|
||||
|
||||
let (mut agent, _tmp) = build_agent_with(provider, vec![], Box::new(NativeToolDispatcher));
|
||||
@@ -738,6 +743,7 @@ async fn turn_handles_none_text_response() {
|
||||
let provider = Box::new(ScriptedProvider::new(vec![ChatResponse {
|
||||
text: None,
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
}]));
|
||||
|
||||
let (mut agent, _tmp) = build_agent_with(provider, vec![], Box::new(NativeToolDispatcher));
|
||||
@@ -761,6 +767,7 @@ async fn turn_preserves_text_alongside_tool_calls() {
|
||||
name: "echo".into(),
|
||||
arguments: r#"{"message": "hi"}"#.into(),
|
||||
}],
|
||||
usage: None,
|
||||
},
|
||||
text_response("Here are the results"),
|
||||
]));
|
||||
@@ -841,6 +848,7 @@ async fn e2e_native_loop_executes_text_fallback_tool_calls_and_persists_history(
|
||||
.into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
},
|
||||
text_response("Completed via tool"),
|
||||
]));
|
||||
@@ -1050,6 +1058,7 @@ async fn native_dispatcher_handles_stringified_arguments() {
|
||||
name: "echo".into(),
|
||||
arguments: r#"{"message": "hello"}"#.into(),
|
||||
}],
|
||||
usage: None,
|
||||
};
|
||||
|
||||
let (_, calls) = dispatcher.parse_response(&response);
|
||||
@@ -1075,6 +1084,7 @@ fn xml_dispatcher_handles_nested_json() {
|
||||
.into(),
|
||||
),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
|
||||
let dispatcher = XmlToolDispatcher;
|
||||
@@ -1092,6 +1102,7 @@ fn xml_dispatcher_handles_empty_tool_call_tag() {
|
||||
let response = ChatResponse {
|
||||
text: Some("<tool_call>\n</tool_call>\nSome text".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
|
||||
let dispatcher = XmlToolDispatcher;
|
||||
@@ -1105,6 +1116,7 @@ fn xml_dispatcher_handles_unclosed_tool_call() {
|
||||
let response = ChatResponse {
|
||||
text: Some("Before\n<tool_call>\n{\"name\": \"shell\"}".into()),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
|
||||
let dispatcher = XmlToolDispatcher;
|
||||
|
||||
@@ -61,6 +61,7 @@ pub struct ToolCall {
|
||||
pub struct ChatResponse {
|
||||
pub text: Option<String>,
|
||||
pub tool_calls: Vec<ToolCall>,
|
||||
pub usage: Option<crate::openhuman::providers::UsageInfo>,
|
||||
}
|
||||
|
||||
impl ChatResponse {
|
||||
@@ -389,6 +390,7 @@ mod tests {
|
||||
let empty = ChatResponse {
|
||||
text: None,
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
assert!(!empty.has_tool_calls());
|
||||
assert_eq!(empty.text_or_empty(), "");
|
||||
@@ -400,6 +402,7 @@ mod tests {
|
||||
name: "shell".into(),
|
||||
arguments: "{}".into(),
|
||||
}],
|
||||
usage: None,
|
||||
};
|
||||
assert!(with_tools.has_tool_calls());
|
||||
assert_eq!(with_tools.text_or_empty(), "Let me check");
|
||||
|
||||
@@ -37,8 +37,21 @@ pub struct AgentConfig {
|
||||
pub max_history_messages: usize,
|
||||
#[serde(default)]
|
||||
pub parallel_tools: bool,
|
||||
/// Maximum number of tool calls to execute concurrently when `parallel_tools` is true.
|
||||
#[serde(default = "default_max_parallel_tools")]
|
||||
pub max_parallel_tools: usize,
|
||||
#[serde(default = "default_agent_tool_dispatcher")]
|
||||
pub tool_dispatcher: String,
|
||||
/// Maximum characters of memory context to inject per turn.
|
||||
/// Higher values provide richer context but consume more of the context window.
|
||||
#[serde(default = "default_max_memory_context_chars")]
|
||||
pub max_memory_context_chars: usize,
|
||||
/// Per-channel maximum permission level for tool execution.
|
||||
/// Keys are channel names (e.g., "telegram", "discord", "web", "cli").
|
||||
/// Values are permission levels: "none", "readonly", "write", "execute", "dangerous".
|
||||
/// Channels not listed default to "readonly".
|
||||
#[serde(default)]
|
||||
pub channel_permissions: std::collections::HashMap<String, String>,
|
||||
}
|
||||
|
||||
fn default_agent_max_tool_iterations() -> usize {
|
||||
@@ -49,10 +62,18 @@ fn default_agent_max_history_messages() -> usize {
|
||||
50
|
||||
}
|
||||
|
||||
fn default_max_parallel_tools() -> usize {
|
||||
4
|
||||
}
|
||||
|
||||
fn default_agent_tool_dispatcher() -> String {
|
||||
"auto".into()
|
||||
}
|
||||
|
||||
fn default_max_memory_context_chars() -> usize {
|
||||
2000
|
||||
}
|
||||
|
||||
impl Default for AgentConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
@@ -60,7 +81,10 @@ impl Default for AgentConfig {
|
||||
max_tool_iterations: default_agent_max_tool_iterations(),
|
||||
max_history_messages: default_agent_max_history_messages(),
|
||||
parallel_tools: false,
|
||||
max_parallel_tools: default_max_parallel_tools(),
|
||||
tool_dispatcher: default_agent_tool_dispatcher(),
|
||||
max_memory_context_chars: default_max_memory_context_chars(),
|
||||
channel_permissions: std::collections::HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,20 +6,16 @@ use serde::{Deserialize, Serialize};
|
||||
/// Which LLM to use for reflection inference.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[derive(Default)]
|
||||
pub enum ReflectionSource {
|
||||
/// Use the local Ollama model via `LocalAiService::prompt()`.
|
||||
/// Model is determined by `config.local_ai.chat_model_id`.
|
||||
#[default]
|
||||
Local,
|
||||
/// Use the cloud reasoning model via `Provider::simple_chat("hint:reasoning")`.
|
||||
Cloud,
|
||||
}
|
||||
|
||||
impl Default for ReflectionSource {
|
||||
fn default() -> Self {
|
||||
Self::Local
|
||||
}
|
||||
}
|
||||
|
||||
/// Configuration for the agent self-learning subsystem.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
pub struct LearningConfig {
|
||||
|
||||
@@ -812,7 +812,7 @@ fn handle_local_ai_set_ollama_path(params: Map<String, Value>) -> ControllerFutu
|
||||
});
|
||||
|
||||
let current_status =
|
||||
serde_json::to_value(&service.status()).map_err(|e| format!("serialize: {e}"))?;
|
||||
serde_json::to_value(service.status()).map_err(|e| format!("serialize: {e}"))?;
|
||||
Ok(serde_json::json!({
|
||||
"ollama_binary_path": new_value,
|
||||
"status": current_status,
|
||||
|
||||
@@ -987,7 +987,11 @@ impl OpenAiCompatibleProvider {
|
||||
}
|
||||
}
|
||||
|
||||
ProviderChatResponse { text, tool_calls }
|
||||
ProviderChatResponse {
|
||||
text,
|
||||
tool_calls,
|
||||
usage: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_native_tool_schema_unsupported(status: reqwest::StatusCode, error: &str) -> bool {
|
||||
@@ -1314,6 +1318,7 @@ impl Provider for OpenAiCompatibleProvider {
|
||||
return Ok(ProviderChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
});
|
||||
}
|
||||
};
|
||||
@@ -1348,7 +1353,11 @@ impl Provider for OpenAiCompatibleProvider {
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
Ok(ProviderChatResponse { text, tool_calls })
|
||||
Ok(ProviderChatResponse {
|
||||
text,
|
||||
tool_calls,
|
||||
usage: None,
|
||||
})
|
||||
}
|
||||
|
||||
async fn chat(
|
||||
@@ -1398,6 +1407,7 @@ impl Provider for OpenAiCompatibleProvider {
|
||||
.map(|text| ProviderChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
})
|
||||
.map_err(|responses_err| {
|
||||
let fb = super::format_anyhow_chain(&responses_err);
|
||||
@@ -1426,6 +1436,7 @@ impl Provider for OpenAiCompatibleProvider {
|
||||
return Ok(ProviderChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -1436,6 +1447,7 @@ impl Provider for OpenAiCompatibleProvider {
|
||||
.map(|text| ProviderChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
})
|
||||
.map_err(|responses_err| {
|
||||
let fb = super::format_anyhow_chain(&responses_err);
|
||||
|
||||
@@ -8,7 +8,7 @@ pub mod traits;
|
||||
#[allow(unused_imports)]
|
||||
pub use traits::{
|
||||
ChatMessage, ChatRequest, ChatResponse, ConversationMessage, Provider, ProviderCapabilityError,
|
||||
ToolCall, ToolResultMessage,
|
||||
ToolCall, ToolResultMessage, UsageInfo,
|
||||
};
|
||||
|
||||
pub use ops::*;
|
||||
|
||||
@@ -49,6 +49,17 @@ pub struct ToolCall {
|
||||
pub arguments: String,
|
||||
}
|
||||
|
||||
/// Token usage information returned by the provider after an inference call.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct UsageInfo {
|
||||
/// Number of tokens in the input/prompt.
|
||||
pub input_tokens: u64,
|
||||
/// Number of tokens in the output/completion.
|
||||
pub output_tokens: u64,
|
||||
/// Total context window size for the model (0 if unknown).
|
||||
pub context_window: u64,
|
||||
}
|
||||
|
||||
/// An LLM response that may contain text, tool calls, or both.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ChatResponse {
|
||||
@@ -56,6 +67,8 @@ pub struct ChatResponse {
|
||||
pub text: Option<String>,
|
||||
/// Tool calls requested by the LLM.
|
||||
pub tool_calls: Vec<ToolCall>,
|
||||
/// Token usage info from the provider (if available).
|
||||
pub usage: Option<UsageInfo>,
|
||||
}
|
||||
|
||||
impl ChatResponse {
|
||||
@@ -75,6 +88,11 @@ impl ChatResponse {
|
||||
pub struct ChatRequest<'a> {
|
||||
pub messages: &'a [ChatMessage],
|
||||
pub tools: Option<&'a [ToolSpec]>,
|
||||
/// Byte offset in the system prompt where static (cacheable) content ends
|
||||
/// and dynamic content begins. Providers that support prompt caching can
|
||||
/// apply `cache_control` to the prefix before this boundary.
|
||||
/// `None` means no cache boundary is known.
|
||||
pub system_prompt_cache_boundary: Option<usize>,
|
||||
}
|
||||
|
||||
/// A tool result to feed back to the LLM.
|
||||
@@ -372,6 +390,7 @@ pub trait Provider: Send + Sync {
|
||||
return Ok(ChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: Vec::new(),
|
||||
usage: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -389,6 +408,7 @@ pub trait Provider: Send + Sync {
|
||||
Ok(ChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: Vec::new(),
|
||||
usage: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -422,6 +442,7 @@ pub trait Provider: Send + Sync {
|
||||
Ok(ChatResponse {
|
||||
text: Some(text),
|
||||
tool_calls: Vec::new(),
|
||||
usage: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -545,6 +566,7 @@ mod tests {
|
||||
let empty = ChatResponse {
|
||||
text: None,
|
||||
tool_calls: vec![],
|
||||
usage: None,
|
||||
};
|
||||
assert!(!empty.has_tool_calls());
|
||||
assert_eq!(empty.text_or_empty(), "");
|
||||
@@ -556,6 +578,7 @@ mod tests {
|
||||
name: "shell".into(),
|
||||
arguments: "{}".into(),
|
||||
}],
|
||||
usage: None,
|
||||
};
|
||||
assert!(with_tools.has_tool_calls());
|
||||
assert_eq!(with_tools.text_or_empty(), "Let me check");
|
||||
@@ -766,6 +789,7 @@ mod tests {
|
||||
let request = ChatRequest {
|
||||
messages: &[ChatMessage::user("Hello")],
|
||||
tools: Some(&tools),
|
||||
system_prompt_cache_boundary: None,
|
||||
};
|
||||
|
||||
let response = provider.chat(request, "model", 0.7).await.unwrap();
|
||||
@@ -783,6 +807,7 @@ mod tests {
|
||||
let request = ChatRequest {
|
||||
messages: &[ChatMessage::user("Hello")],
|
||||
tools: None,
|
||||
system_prompt_cache_boundary: None,
|
||||
};
|
||||
|
||||
let response = provider.chat(request, "model", 0.7).await.unwrap();
|
||||
@@ -883,6 +908,7 @@ mod tests {
|
||||
ChatMessage::system("BASE_SYSTEM_PROMPT"),
|
||||
],
|
||||
tools: Some(&tools),
|
||||
system_prompt_cache_boundary: None,
|
||||
};
|
||||
|
||||
let response = provider.chat(request, "model", 0.7).await.unwrap();
|
||||
@@ -905,6 +931,7 @@ mod tests {
|
||||
let request = ChatRequest {
|
||||
messages: &[ChatMessage::system("BASE"), ChatMessage::user("Hello")],
|
||||
tools: Some(&tools),
|
||||
system_prompt_cache_boundary: None,
|
||||
};
|
||||
|
||||
let response = provider.chat(request, "model", 0.7).await.unwrap();
|
||||
@@ -927,6 +954,7 @@ mod tests {
|
||||
let request = ChatRequest {
|
||||
messages: &[ChatMessage::user("Hello")],
|
||||
tools: Some(&tools),
|
||||
system_prompt_cache_boundary: None,
|
||||
};
|
||||
|
||||
let err = provider.chat(request, "model", 0.7).await.unwrap_err();
|
||||
|
||||
@@ -35,7 +35,7 @@ pub fn register<'js>(
|
||||
|
||||
router
|
||||
.register(&tunnel_uuid, &sc.skill_id, name, backend_id)
|
||||
.map_err(|e| js_err(e))
|
||||
.map_err(js_err)
|
||||
},
|
||||
),
|
||||
)?;
|
||||
@@ -56,7 +56,7 @@ pub fn register<'js>(
|
||||
|
||||
router
|
||||
.unregister(&tunnel_uuid, &sc.skill_id)
|
||||
.map_err(|e| js_err(e))
|
||||
.map_err(js_err)
|
||||
},
|
||||
),
|
||||
)?;
|
||||
|
||||
@@ -67,7 +67,7 @@ pub use schemas::{
|
||||
pub use screenshot::ScreenshotTool;
|
||||
pub use shell::ShellTool;
|
||||
pub use tool_stats::ToolStatsTool;
|
||||
pub use traits::Tool;
|
||||
pub use traits::{PermissionLevel, Tool};
|
||||
#[allow(unused_imports)]
|
||||
pub use traits::{ToolResult, ToolSpec};
|
||||
pub use web_search_tool::WebSearchTool;
|
||||
|
||||
@@ -1,6 +1,38 @@
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Permission level required to execute a tool.
|
||||
///
|
||||
/// Channels can set a maximum permission level to restrict which tools
|
||||
/// are available. Tools requiring a level above the channel's maximum
|
||||
/// are rejected before execution.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default)]
|
||||
pub enum PermissionLevel {
|
||||
/// No permission needed (metadata-only operations).
|
||||
None = 0,
|
||||
/// Read-only operations (file reads, memory recall, listing).
|
||||
#[default]
|
||||
ReadOnly = 1,
|
||||
/// Write operations (file writes, memory store).
|
||||
Write = 2,
|
||||
/// Command execution (shell, scripts).
|
||||
Execute = 3,
|
||||
/// Dangerous/destructive operations (hardware, system-level).
|
||||
Dangerous = 4,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for PermissionLevel {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::None => write!(f, "None"),
|
||||
Self::ReadOnly => write!(f, "ReadOnly"),
|
||||
Self::Write => write!(f, "Write"),
|
||||
Self::Execute => write!(f, "Execute"),
|
||||
Self::Dangerous => write!(f, "Dangerous"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Result of a tool execution
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolResult {
|
||||
@@ -32,6 +64,13 @@ pub trait Tool: Send + Sync {
|
||||
/// Execute the tool with given arguments
|
||||
async fn execute(&self, args: serde_json::Value) -> anyhow::Result<ToolResult>;
|
||||
|
||||
/// Permission level required to execute this tool.
|
||||
/// Channels with a lower maximum permission level will reject this tool.
|
||||
/// Default: `ReadOnly`. Override for write/execute/dangerous tools.
|
||||
fn permission_level(&self) -> PermissionLevel {
|
||||
PermissionLevel::ReadOnly
|
||||
}
|
||||
|
||||
/// Get the full spec for LLM registration
|
||||
fn spec(&self) -> ToolSpec {
|
||||
ToolSpec {
|
||||
|
||||
Reference in New Issue
Block a user