conversation_api/execution/
llm.rs1use crate::execution::{ExternalError, InvocationContext};
3use async_trait::async_trait;
4pub use llm_api::{
5 Completion, CompletionRequest, ContentPart, Continuation, FinishReason, INPUT_AUDIO,
6 INPUT_FILE, INPUT_IMAGE, INPUT_VIDEO, Message, MessageRole, ModelCapabilities,
7 ModelConstraints, ModelMode, ModelProfile, TokenUsage, ToolCall, UseCase,
8 input_modality_for_kind, payload_input_modalities,
9};
10use serde::{Deserialize, Serialize};
11use serde_json::Value;
12use thiserror::Error;
13
14#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
16pub struct ToolDefinition {
17 pub name: String,
19 pub description: String,
21 pub input_schema: Value,
23 pub policy: ToolPolicy,
25}
26
27#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
29#[serde(rename_all = "snake_case")]
30pub enum ToolEffect {
31 ReadOnly,
33 Mutating,
35}
36
37#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
39#[serde(rename_all = "snake_case")]
40pub enum ApprovalRequirement {
41 Never,
43 WhenInteractive,
45 Always,
47}
48
49impl ApprovalRequirement {
50 #[must_use]
52 pub const fn rank(self) -> u8 {
53 match self {
54 Self::Never => 0,
55 Self::WhenInteractive => 1,
56 Self::Always => 2,
57 }
58 }
59}
60
61#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
63#[serde(tag = "type", rename_all = "snake_case")]
64pub enum ApprovalPolicy {
65 Fixed {
67 requirement: ApprovalRequirement,
69 },
70 PerInvocation {
72 minimum: ApprovalRequirement,
74 },
75}
76
77impl ApprovalPolicy {
78 #[must_use]
80 pub const fn permits(self, requirement: ApprovalRequirement) -> bool {
81 match self {
82 Self::Fixed { requirement: fixed } => fixed.rank() == requirement.rank(),
83 Self::PerInvocation { minimum } => requirement.rank() >= minimum.rank(),
84 }
85 }
86
87 #[must_use]
89 pub const fn is_fixed_never(self) -> bool {
90 matches!(
91 self,
92 Self::Fixed {
93 requirement: ApprovalRequirement::Never
94 }
95 )
96 }
97}
98
99#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
101pub struct ToolPolicy {
102 pub effect: ToolEffect,
104 pub approval: ApprovalPolicy,
106 pub parallel_safe: bool,
108}
109
110#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
112#[serde(deny_unknown_fields)]
113pub struct UsageSummary {
114 pub model_requests: u32,
116 pub reported_model_requests: u32,
118 pub input_tokens: u64,
120 pub output_tokens: u64,
122 pub cached_input_tokens: Option<u64>,
124 pub reasoning_output_tokens: Option<u64>,
126 pub peak_input_tokens: u64,
128 pub credits: Option<u64>,
130 pub tool_calls: u32,
132 pub window_truncations: u32,
134 pub trimmed_messages: u64,
136}
137
138impl UsageSummary {
139 pub fn record_model_attempt(&mut self) {
141 self.model_requests = self.model_requests.saturating_add(1);
142 }
143
144 pub fn record_usage(&mut self, usage: TokenUsage, reported: bool) {
146 if reported {
147 self.reported_model_requests = self.reported_model_requests.saturating_add(1);
148 }
149 self.input_tokens = self.input_tokens.saturating_add(usage.input_tokens);
150 self.output_tokens = self.output_tokens.saturating_add(usage.output_tokens);
151 self.peak_input_tokens = self.peak_input_tokens.max(usage.input_tokens);
152 add_optional(&mut self.cached_input_tokens, usage.cached_input_tokens);
153 add_optional(
154 &mut self.reasoning_output_tokens,
155 usage.reasoning_output_tokens,
156 );
157 add_optional(&mut self.credits, usage.credits);
158 }
159
160 pub fn record_model_request(&mut self, usage: Option<TokenUsage>) {
162 self.record_model_attempt();
163 let Some(usage) = usage else {
164 return;
165 };
166 self.record_usage(usage, true);
167 }
168
169 pub fn record_window_truncation(&mut self, removed_messages: usize) {
171 self.window_truncations = self.window_truncations.saturating_add(1);
172 self.trimmed_messages = self
173 .trimmed_messages
174 .saturating_add(u64::try_from(removed_messages).unwrap_or(u64::MAX));
175 }
176}
177
178fn add_optional(total: &mut Option<u64>, value: Option<u64>) {
179 if let Some(value) = value {
180 *total = Some(total.unwrap_or(0).saturating_add(value));
181 }
182}
183
184#[derive(Clone, Debug, Error, Eq, PartialEq)]
186#[error("{error}")]
187pub struct LlmFailure {
188 pub error: ExternalError,
190 pub usage: Option<TokenUsage>,
192}
193
194impl From<ExternalError> for LlmFailure {
195 fn from(error: ExternalError) -> Self {
196 Self { error, usage: None }
197 }
198}
199
200#[async_trait]
205pub trait Llm: Send + Sync {
206 async fn model_profile(
208 &self,
209 context: &InvocationContext,
210 use_case: &UseCase,
211 model_mode: &ModelMode,
212 ) -> Result<ModelProfile, ExternalError>;
213
214 async fn complete(
216 &self,
217 context: &InvocationContext,
218 request: CompletionRequest,
219 ) -> Result<Completion, LlmFailure>;
220}
221
222impl ToolDefinition {
223 pub fn model_definition(&self) -> llm_api::ToolDefinition {
224 llm_api::ToolDefinition {
225 name: self.name.clone(),
226 description: self.description.clone(),
227 input_schema: self.input_schema.clone(),
228 }
229 }
230}