use crate::execution::{ExternalError, InvocationContext};
use async_trait::async_trait;
pub use llm_api::{
Completion, CompletionRequest, ContentPart, Continuation, FinishReason, Message, MessageRole,
ModelConstraints, ModelMode, ModelProfile, TokenUsage, ToolCall, UseCase,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use thiserror::Error;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ToolDefinition {
pub name: String,
pub description: String,
pub input_schema: Value,
pub policy: ToolPolicy,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ToolEffect {
ReadOnly,
Mutating,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ApprovalRequirement {
Never,
WhenInteractive,
Always,
}
impl ApprovalRequirement {
#[must_use]
pub const fn rank(self) -> u8 {
match self {
Self::Never => 0,
Self::WhenInteractive => 1,
Self::Always => 2,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ApprovalPolicy {
Fixed {
requirement: ApprovalRequirement,
},
PerInvocation {
minimum: ApprovalRequirement,
},
}
impl ApprovalPolicy {
#[must_use]
pub const fn permits(self, requirement: ApprovalRequirement) -> bool {
match self {
Self::Fixed { requirement: fixed } => fixed.rank() == requirement.rank(),
Self::PerInvocation { minimum } => requirement.rank() >= minimum.rank(),
}
}
#[must_use]
pub const fn is_fixed_never(self) -> bool {
matches!(
self,
Self::Fixed {
requirement: ApprovalRequirement::Never
}
)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub struct ToolPolicy {
pub effect: ToolEffect,
pub approval: ApprovalPolicy,
pub parallel_safe: bool,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
pub struct UsageSummary {
pub model_requests: u32,
pub reported_model_requests: u32,
pub input_tokens: u64,
pub output_tokens: u64,
pub cached_input_tokens: Option<u64>,
pub reasoning_output_tokens: Option<u64>,
pub peak_input_tokens: u64,
pub credits: Option<u64>,
pub tool_calls: u32,
pub window_truncations: u32,
pub trimmed_messages: u64,
}
impl UsageSummary {
pub fn record_model_attempt(&mut self) {
self.model_requests = self.model_requests.saturating_add(1);
}
pub fn record_usage(&mut self, usage: TokenUsage, reported: bool) {
if reported {
self.reported_model_requests = self.reported_model_requests.saturating_add(1);
}
self.input_tokens = self.input_tokens.saturating_add(usage.input_tokens);
self.output_tokens = self.output_tokens.saturating_add(usage.output_tokens);
self.peak_input_tokens = self.peak_input_tokens.max(usage.input_tokens);
add_optional(&mut self.cached_input_tokens, usage.cached_input_tokens);
add_optional(
&mut self.reasoning_output_tokens,
usage.reasoning_output_tokens,
);
add_optional(&mut self.credits, usage.credits);
}
pub fn record_model_request(&mut self, usage: Option<TokenUsage>) {
self.record_model_attempt();
let Some(usage) = usage else {
return;
};
self.record_usage(usage, true);
}
pub fn record_window_truncation(&mut self, removed_messages: usize) {
self.window_truncations = self.window_truncations.saturating_add(1);
self.trimmed_messages = self
.trimmed_messages
.saturating_add(u64::try_from(removed_messages).unwrap_or(u64::MAX));
}
}
fn add_optional(total: &mut Option<u64>, value: Option<u64>) {
if let Some(value) = value {
*total = Some(total.unwrap_or(0).saturating_add(value));
}
}
#[derive(Clone, Debug, Error, Eq, PartialEq)]
#[error("{error}")]
pub struct LlmFailure {
pub error: ExternalError,
pub usage: Option<TokenUsage>,
}
impl From<ExternalError> for LlmFailure {
fn from(error: ExternalError) -> Self {
Self { error, usage: None }
}
}
#[async_trait]
pub trait Llm: Send + Sync {
async fn model_profile(
&self,
context: &InvocationContext,
use_case: &UseCase,
model_mode: &ModelMode,
) -> Result<ModelProfile, ExternalError>;
async fn complete(
&self,
context: &InvocationContext,
request: CompletionRequest,
) -> Result<Completion, LlmFailure>;
}
impl ToolDefinition {
pub fn model_definition(&self) -> llm_api::ToolDefinition {
llm_api::ToolDefinition {
name: self.name.clone(),
description: self.description.clone(),
input_schema: self.input_schema.clone(),
}
}
}