1use std::pin::Pin;
2
3use futures::Stream;
4use serde::{Deserialize, Serialize};
5
6use crate::extension::InferenceEngineId;
7use crate::lifecycle::TurnCleanupOwnership;
8use crate::reliability::ReliabilityRequestPolicy;
9use crate::tools::{ToolChoice, ToolSpec};
10use crate::transcript::TranscriptItem;
11
12#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
13pub struct ModelSelection {
14 pub provider: String,
15 pub model: String,
16}
17
18#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
19#[serde(rename_all = "snake_case")]
20pub enum ProviderAuthType {
21 None,
22 ApiKey,
23 OAuth,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
27pub struct InferenceProviderMetadata {
28 pub name: String,
29 pub description: Option<String>,
30 pub auth_type: ProviderAuthType,
31 pub auth_label: Option<String>,
32 pub auth_configured: Option<bool>,
33 pub recommended: bool,
34 pub sort_order: i32,
35}
36
37impl InferenceProviderMetadata {
38 pub fn local(name: impl Into<String>) -> Self {
39 Self {
40 name: name.into(),
41 description: None,
42 auth_type: ProviderAuthType::None,
43 auth_label: None,
44 auth_configured: Some(true),
45 recommended: false,
46 sort_order: 100,
47 }
48 }
49}
50
51#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
52#[serde(rename_all = "snake_case")]
53pub enum ToolSearchMode {
54 #[default]
55 Explicit,
56 Auto,
57 ProviderNative,
58}
59
60impl ToolSearchMode {
61 pub fn allows_provider_native(self) -> bool {
62 matches!(self, Self::Auto | Self::ProviderNative)
63 }
64}
65
66#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
67#[serde(rename_all = "snake_case")]
68pub enum ToolSearchProviderVariant {
69 #[default]
70 Default,
71 Regex,
72 Bm25,
73}
74
75#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
76#[serde(rename_all = "camelCase")]
77pub struct ToolSearchConfig {
78 #[serde(default)]
79 pub mode: ToolSearchMode,
80 #[serde(default, skip_serializing_if = "Option::is_none")]
81 pub max_catalog_items: Option<u32>,
82 #[serde(default)]
83 pub include_mcp: bool,
84 #[serde(default)]
85 pub include_skills: bool,
86 #[serde(default)]
87 pub fallback_to_explicit_tools: bool,
88 #[serde(default)]
89 pub provider_variant: ToolSearchProviderVariant,
90}
91
92impl Default for ToolSearchConfig {
93 fn default() -> Self {
94 Self {
95 mode: ToolSearchMode::Explicit,
96 max_catalog_items: None,
97 include_mcp: true,
98 include_skills: true,
99 fallback_to_explicit_tools: true,
100 provider_variant: ToolSearchProviderVariant::Default,
101 }
102 }
103}
104
105impl ToolSearchConfig {
106 pub fn explicit() -> Self {
107 Self {
108 mode: ToolSearchMode::Explicit,
109 ..Self::default()
110 }
111 }
112
113 pub fn provider_native() -> Self {
114 Self {
115 mode: ToolSearchMode::ProviderNative,
116 ..Self::default()
117 }
118 }
119
120 pub fn is_provider_native_requested(&self) -> bool {
121 self.mode.allows_provider_native()
122 }
123
124 pub fn resolve_effective_mode(
133 &self,
134 provider_native_supported: bool,
135 ) -> Result<EffectiveToolSearchMode, ToolSearchModeError> {
136 match self.mode {
137 ToolSearchMode::Explicit => Ok(EffectiveToolSearchMode::Explicit),
138 ToolSearchMode::Auto => {
139 if provider_native_supported {
140 Ok(EffectiveToolSearchMode::ProviderNative)
141 } else {
142 Ok(EffectiveToolSearchMode::Explicit)
143 }
144 }
145 ToolSearchMode::ProviderNative => {
146 if provider_native_supported {
147 Ok(EffectiveToolSearchMode::ProviderNative)
148 } else if self.fallback_to_explicit_tools {
149 Ok(EffectiveToolSearchMode::Explicit)
150 } else {
151 Err(ToolSearchModeError::ProviderNativeUnsupported)
152 }
153 }
154 }
155 }
156}
157
158#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
159#[serde(rename_all = "snake_case")]
160pub enum EffectiveToolSearchMode {
161 Explicit,
162 ProviderNative,
163}
164
165#[derive(Debug, Clone, Copy, PartialEq, Eq)]
166pub enum ToolSearchModeError {
167 ProviderNativeUnsupported,
168}
169
170impl std::fmt::Display for ToolSearchModeError {
171 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
172 match self {
173 Self::ProviderNativeUnsupported => write!(
174 f,
175 "provider-native tool search was requested but the selected provider/model does \
176 not support it and fallback_to_explicit_tools is disabled; enable fallback or \
177 pick a supported model"
178 ),
179 }
180 }
181}
182
183impl std::error::Error for ToolSearchModeError {}
184
185#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
186#[serde(rename_all = "camelCase")]
187pub struct ToolSearchConfigOverlay {
188 #[serde(default, skip_serializing_if = "Option::is_none")]
189 pub mode: Option<ToolSearchMode>,
190 #[serde(default, skip_serializing_if = "Option::is_none")]
191 pub max_catalog_items: Option<u32>,
192 #[serde(default, skip_serializing_if = "Option::is_none")]
193 pub include_mcp: Option<bool>,
194 #[serde(default, skip_serializing_if = "Option::is_none")]
195 pub include_skills: Option<bool>,
196 #[serde(default, skip_serializing_if = "Option::is_none")]
197 pub fallback_to_explicit_tools: Option<bool>,
198 #[serde(default, skip_serializing_if = "Option::is_none")]
199 pub provider_variant: Option<ToolSearchProviderVariant>,
200}
201
202impl ToolSearchConfigOverlay {
203 pub fn overlay(&mut self, other: &Self) {
204 if other.mode.is_some() {
205 self.mode = other.mode;
206 }
207 if other.max_catalog_items.is_some() {
208 self.max_catalog_items = other.max_catalog_items;
209 }
210 if other.include_mcp.is_some() {
211 self.include_mcp = other.include_mcp;
212 }
213 if other.include_skills.is_some() {
214 self.include_skills = other.include_skills;
215 }
216 if other.fallback_to_explicit_tools.is_some() {
217 self.fallback_to_explicit_tools = other.fallback_to_explicit_tools;
218 }
219 if other.provider_variant.is_some() {
220 self.provider_variant = other.provider_variant;
221 }
222 }
223
224 pub fn apply_to(&self, config: &mut ToolSearchConfig) {
225 if let Some(mode) = self.mode {
226 config.mode = mode;
227 }
228 if let Some(max_catalog_items) = self.max_catalog_items {
229 config.max_catalog_items = Some(max_catalog_items);
230 }
231 if let Some(include_mcp) = self.include_mcp {
232 config.include_mcp = include_mcp;
233 }
234 if let Some(include_skills) = self.include_skills {
235 config.include_skills = include_skills;
236 }
237 if let Some(fallback_to_explicit_tools) = self.fallback_to_explicit_tools {
238 config.fallback_to_explicit_tools = fallback_to_explicit_tools;
239 }
240 if let Some(provider_variant) = self.provider_variant {
241 config.provider_variant = provider_variant;
242 }
243 }
244}
245
246#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
247pub struct InstructionBundle {
248 pub system: Option<String>,
249 pub developer: Option<String>,
250 pub developer_context: Option<String>,
257}
258
259#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
260#[serde(rename_all = "snake_case")]
261pub enum RuntimeProfile {
262 #[default]
263 Interactive,
264 NonInteractive,
265 Eval,
266}
267
268impl RuntimeProfile {
269 pub fn as_str(self) -> &'static str {
270 match self {
271 Self::Interactive => "interactive",
272 Self::NonInteractive => "non_interactive",
273 Self::Eval => "eval",
274 }
275 }
276
277 pub fn is_non_interactive(self) -> bool {
278 matches!(self, Self::NonInteractive | Self::Eval)
279 }
280}
281
282impl std::str::FromStr for RuntimeProfile {
283 type Err = anyhow::Error;
284
285 fn from_str(value: &str) -> Result<Self, Self::Err> {
286 match value.trim().to_ascii_lowercase().as_str() {
287 "interactive" => Ok(Self::Interactive),
288 "non_interactive" | "non-interactive" | "headless" => Ok(Self::NonInteractive),
289 "eval" => Ok(Self::Eval),
290 other => anyhow::bail!(
291 "unsupported runtime profile {other:?}; expected interactive, non_interactive, or eval"
292 ),
293 }
294 }
295}
296
297#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
298pub struct ReasoningConfig {
299 pub enabled: bool,
300 pub level: Option<String>,
301}
302
303#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
304#[serde(rename_all = "snake_case")]
305pub enum ProviderFamily {
306 #[default]
307 Mock,
308 OpenAi,
309 Anthropic,
310 Gemini,
311 Xai,
312 Opencode,
313 Poolside,
314 Cursor,
315}
316
317#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
318#[serde(rename_all = "snake_case")]
319pub enum ModelSchemaPolicy {
320 #[default]
321 StandardRequiredFirst,
322 RequiredFirstFlat,
323}
324
325#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
326#[serde(rename_all = "snake_case")]
327pub enum ModelInstructionOverlay {
328 #[default]
329 Standard,
330 LiteralToolOutputs,
331 IntuitiveContext,
332}
333
334#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
335#[serde(rename_all = "camelCase")]
336pub struct ModelProfileReasoning {
337 #[serde(default, skip_serializing_if = "Option::is_none")]
338 pub orientation: Option<String>,
339 #[serde(default, skip_serializing_if = "Option::is_none")]
340 pub execution: Option<String>,
341 #[serde(default, skip_serializing_if = "Option::is_none")]
342 pub verification: Option<String>,
343 #[serde(default, skip_serializing_if = "Option::is_none")]
344 pub recovery: Option<String>,
345}
346
347#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
348#[serde(rename_all = "camelCase")]
349pub struct ModelHarnessProfile {
350 pub model: String,
351 pub provider: String,
352 pub provider_family: ProviderFamily,
353 #[serde(default, skip_serializing_if = "Option::is_none")]
354 pub edit_tool: Option<String>,
355 #[serde(default)]
356 pub schema_policy: ModelSchemaPolicy,
357 #[serde(default)]
358 pub instruction_overlay: ModelInstructionOverlay,
359 #[serde(default)]
360 pub reasoning: ModelProfileReasoning,
361 #[serde(default, skip_serializing_if = "Option::is_none")]
362 pub parallel_tool_calls: Option<bool>,
363 #[serde(default, skip_serializing_if = "Option::is_none")]
364 pub auto_compact_token_limit: Option<u32>,
365}
366
367#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
368#[serde(rename_all = "snake_case")]
369pub enum SpeedPolicyPhase {
370 #[default]
371 Orientation,
372 Execution,
373 Verification,
374 Recovery,
375}
376
377impl SpeedPolicyPhase {
378 pub fn as_str(self) -> &'static str {
379 match self {
380 Self::Orientation => "orientation",
381 Self::Execution => "execution",
382 Self::Verification => "verification",
383 Self::Recovery => "recovery",
384 }
385 }
386}
387
388#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
389#[serde(rename_all = "camelCase")]
390pub struct SpeedPolicyDecision {
391 pub phase: SpeedPolicyPhase,
392 pub desired_reasoning: String,
393 pub applied_reasoning: Option<String>,
394 pub supported: bool,
395}
396
397#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
398pub struct OutputConfig {
399 pub max_tokens: Option<u32>,
400 pub temperature: Option<f32>,
401 pub top_p: Option<f32>,
402 pub response_format: Option<serde_json::Value>,
403}
404
405#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
406#[serde(rename_all = "snake_case")]
407pub enum HostedWebSearchMode {
408 #[default]
409 Disabled,
410 Cached,
411 Live,
412}
413
414#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
415pub struct HostedWebSearchConfig {
416 pub mode: HostedWebSearchMode,
417}
418
419impl HostedWebSearchConfig {
420 pub fn disabled() -> Self {
421 Self {
422 mode: HostedWebSearchMode::Disabled,
423 }
424 }
425
426 pub fn cached() -> Self {
427 Self {
428 mode: HostedWebSearchMode::Cached,
429 }
430 }
431
432 pub fn live() -> Self {
433 Self {
434 mode: HostedWebSearchMode::Live,
435 }
436 }
437
438 pub fn is_enabled(&self) -> bool {
439 self.mode != HostedWebSearchMode::Disabled
440 }
441}
442
443#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
444pub struct RuntimeHints {
445 pub trace_id: Option<String>,
446 pub prompt_cache_key: Option<String>,
447 pub auto_compact_token_limit: Option<u32>,
448 #[serde(default)]
449 pub profile: RuntimeProfile,
450 #[serde(default, skip_serializing_if = "Option::is_none")]
451 pub parallel_tool_calls: Option<bool>,
452 #[serde(default)]
453 pub hosted_web_search: HostedWebSearchConfig,
454 #[serde(default)]
455 pub tool_search: ToolSearchConfig,
456 #[serde(default, skip_serializing_if = "Option::is_none")]
457 pub speed_policy: Option<SpeedPolicyDecision>,
458 #[serde(default, skip_serializing_if = "Option::is_none")]
459 pub deadline_remaining_seconds: Option<u64>,
460 #[serde(default, skip_serializing_if = "Option::is_none")]
461 pub reliability: Option<ReliabilityRequestPolicy>,
462 #[serde(default, skip_serializing_if = "Option::is_none")]
470 pub service_tier: Option<String>,
471}
472
473#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
474pub struct AgentInferenceRequest {
475 pub model: ModelSelection,
476 pub instructions: InstructionBundle,
477 pub transcript: Vec<TranscriptItem>,
478 pub tools: Vec<ToolSpec>,
479 pub tool_choice: ToolChoice,
480 pub reasoning: ReasoningConfig,
481 pub output: OutputConfig,
482 pub runtime: RuntimeHints,
483 pub metadata: serde_json::Value,
484}
485
486#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
487pub struct MessageDelta {
488 pub text: String,
489 #[serde(default, skip_serializing_if = "Option::is_none")]
490 pub phase: Option<String>,
491}
492
493#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
494pub struct ReasoningDelta {
495 pub text: String,
496}
497
498#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
499pub struct ToolCallStarted {
500 pub id: String,
501 pub name: String,
502}
503
504#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
505pub struct ToolCallDelta {
506 pub id: String,
507 pub arguments_delta: String,
508}
509
510#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
511pub struct ToolCallCompleted {
512 pub id: String,
513 pub name: String,
514 pub arguments: String,
515}
516
517#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
518pub struct HostedToolCallStarted {
519 pub id: String,
520 pub name: String,
521}
522
523#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
524pub struct HostedToolCallCompleted {
525 pub id: String,
526 pub name: String,
527 pub arguments: String,
528}
529
530#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
531pub struct TokenUsage {
532 pub prompt_tokens: u32,
533 pub completion_tokens: u32,
534 pub total_tokens: u32,
535 #[serde(default)]
536 pub cached_prompt_tokens: u32,
537 #[serde(default)]
544 pub cache_creation_prompt_tokens: u32,
545 #[serde(default, skip_serializing_if = "Option::is_none")]
546 pub cache_hit_rate: Option<f64>,
547 #[serde(default, skip_serializing_if = "Option::is_none")]
556 pub service_tier: Option<String>,
557}
558
559impl TokenUsage {
560 pub fn new(prompt_tokens: u32, completion_tokens: u32, total_tokens: u32) -> Self {
561 Self {
562 prompt_tokens,
563 completion_tokens,
564 total_tokens,
565 cached_prompt_tokens: 0,
566 cache_creation_prompt_tokens: 0,
567 cache_hit_rate: cache_hit_rate(prompt_tokens, 0),
568 service_tier: None,
569 }
570 }
571
572 pub fn with_service_tier(mut self, service_tier: Option<String>) -> Self {
573 self.service_tier = service_tier;
574 self
575 }
576
577 pub fn with_cached_prompt_tokens(mut self, cached_prompt_tokens: u32) -> Self {
578 self.cached_prompt_tokens = cached_prompt_tokens.min(self.prompt_tokens);
579 self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
580 self
581 }
582
583 pub fn with_cache_creation_prompt_tokens(mut self, cache_creation_prompt_tokens: u32) -> Self {
584 self.cache_creation_prompt_tokens = cache_creation_prompt_tokens.min(self.prompt_tokens);
585 self
586 }
587
588 pub fn add_assign(&mut self, usage: &TokenUsage) {
589 self.prompt_tokens = self.prompt_tokens.saturating_add(usage.prompt_tokens);
590 self.completion_tokens = self
591 .completion_tokens
592 .saturating_add(usage.completion_tokens);
593 self.total_tokens = self.total_tokens.saturating_add(usage.total_tokens);
594 self.cached_prompt_tokens = self
595 .cached_prompt_tokens
596 .saturating_add(usage.cached_prompt_tokens);
597 self.cache_creation_prompt_tokens = self
598 .cache_creation_prompt_tokens
599 .saturating_add(usage.cache_creation_prompt_tokens);
600 self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
601 if usage.service_tier.is_some() {
602 self.service_tier.clone_from(&usage.service_tier);
603 }
604 }
605
606 pub fn is_empty(&self) -> bool {
607 self.prompt_tokens == 0
608 && self.completion_tokens == 0
609 && self.total_tokens == 0
610 && self.cached_prompt_tokens == 0
611 && self.cache_creation_prompt_tokens == 0
612 }
613}
614
615pub fn cache_hit_rate(prompt_tokens: u32, cached_prompt_tokens: u32) -> Option<f64> {
616 if prompt_tokens == 0 {
617 None
618 } else {
619 Some(f64::from(cached_prompt_tokens.min(prompt_tokens)) / f64::from(prompt_tokens))
620 }
621}
622
623#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
624pub struct CompletionMetadata {
625 pub stop_reason: Option<String>,
626 pub provider_response_id: Option<String>,
627}
628
629pub fn finish_reason_from_stop_reason(stop_reason: &str) -> String {
637 match stop_reason {
638 "end_turn" | "stop" | "stop_sequence" => "stop",
639 "max_tokens" | "length" => "length",
640 "tool_use" | "tool_calls" => "toolUse",
641 "content_filter" => "contentFilter",
642 "refusal" => "refusal",
643 other => other,
644 }
645 .to_string()
646}
647
648#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
649pub struct InferenceFailure {
650 pub message: String,
651}
652
653#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
654pub struct CompactionProgress {
655 pub status: String,
656 #[serde(default, skip_serializing_if = "Option::is_none")]
657 pub item_id: Option<String>,
658 #[serde(default, skip_serializing_if = "Option::is_none")]
660 pub tokens_before: Option<u32>,
661 #[serde(default, skip_serializing_if = "Option::is_none")]
663 pub tokens_after: Option<u32>,
664 #[serde(default, skip_serializing_if = "Option::is_none")]
666 pub duration_ms: Option<u64>,
667 #[serde(default, skip_serializing_if = "Option::is_none")]
671 pub item: Option<serde_json::Value>,
672}
673
674#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
675pub enum InferenceEvent {
676 MessageDelta(MessageDelta),
677 ReasoningDelta(ReasoningDelta),
678 ToolCallStarted(ToolCallStarted),
679 ToolCallDelta(ToolCallDelta),
680 ToolCallCompleted(ToolCallCompleted),
681 OutputItemCompleted(serde_json::Value),
684 HostedToolCallStarted(HostedToolCallStarted),
685 HostedToolCallCompleted(HostedToolCallCompleted),
686 Compaction(CompactionProgress),
687 Usage(TokenUsage),
688 Completed(CompletionMetadata),
689 Failed(InferenceFailure),
690 ProviderMetadata(serde_json::Value),
691}
692
693pub type InferenceEventStream =
694 Pin<Box<dyn Stream<Item = anyhow::Result<InferenceEvent>> + Send + 'static>>;
695
696#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
697pub struct InferenceCapabilities {
698 pub streaming: bool,
699 pub tool_calls: bool,
700 pub parallel_tool_calls: bool,
701 pub reasoning_summaries: bool,
702 pub structured_output: bool,
703 pub image_input: bool,
704 pub prompt_cache: bool,
705 pub provider_metadata: bool,
706 pub tool_search: bool,
707}
708
709impl InferenceCapabilities {
710 pub fn text_only() -> Self {
711 Self {
712 streaming: true,
713 tool_calls: false,
714 parallel_tool_calls: false,
715 reasoning_summaries: false,
716 structured_output: false,
717 image_input: false,
718 prompt_cache: false,
719 provider_metadata: false,
720 tool_search: false,
721 }
722 }
723
724 pub fn coding_agent_default() -> Self {
725 Self {
726 streaming: true,
727 tool_calls: true,
728 parallel_tool_calls: true,
729 reasoning_summaries: false,
730 structured_output: false,
731 image_input: false,
732 prompt_cache: false,
733 provider_metadata: true,
734 tool_search: false,
735 }
736 }
737}
738
739#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
740pub struct ModelDescriptor {
741 pub id: String,
742 pub name: String,
743 pub context_window: Option<u32>,
744 #[serde(default, skip_serializing_if = "Option::is_none")]
745 pub default_reasoning: Option<String>,
746 #[serde(default, skip_serializing_if = "Vec::is_empty")]
747 pub supported_reasoning: Vec<ReasoningEffortDescriptor>,
748}
749
750#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
751pub struct ReasoningEffortDescriptor {
752 pub effort: String,
753 pub description: String,
754}
755
756pub struct InferenceProviderContext<'a> {
757 pub provider_id: &'a str,
758}
759
760pub struct InferenceTurnContext<'a> {
761 pub thread_id: &'a str,
762 pub turn_id: &'a str,
763 pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
770}
771
772#[derive(Debug, Clone)]
774pub struct TurnToolOutcome {
775 pub result: String,
776 pub is_error: bool,
777}
778
779#[async_trait::async_trait]
783pub trait TurnToolExecutor: Send + Sync {
784 async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
785
786 fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
792}
793
794#[async_trait::async_trait]
797pub trait ProviderTurnCleanup: Send + Sync {
798 fn ownership(&self) -> TurnCleanupOwnership;
800
801 async fn wait_for_cleanup(&self) -> anyhow::Result<()>;
803}
804
805#[async_trait::async_trait]
806pub trait InferenceEngine: Send + Sync + 'static {
807 fn id(&self) -> InferenceEngineId;
808 fn capabilities(&self) -> InferenceCapabilities;
809
810 fn metadata(&self) -> InferenceProviderMetadata {
811 InferenceProviderMetadata::local(self.id())
812 }
813
814 async fn list_models(
815 &self,
816 ctx: InferenceProviderContext<'_>,
817 ) -> anyhow::Result<Vec<ModelDescriptor>>;
818
819 async fn compact_turn(
822 &self,
823 _ctx: InferenceTurnContext<'_>,
824 _request: AgentInferenceRequest,
825 ) -> anyhow::Result<Option<InferenceEventStream>> {
826 Ok(None)
827 }
828
829 async fn stream_turn(
830 &self,
831 ctx: InferenceTurnContext<'_>,
832 request: AgentInferenceRequest,
833 ) -> anyhow::Result<InferenceEventStream>;
834}
835
836#[cfg(test)]
837mod tests {
838 use super::*;
839
840 #[test]
841 fn finish_reason_mapping_normalizes_known_stop_reasons() {
842 assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
843 assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
844 assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
845 assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
846 assert_eq!(finish_reason_from_stop_reason("length"), "length");
847 assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
848 assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
849 assert_eq!(
850 finish_reason_from_stop_reason("content_filter"),
851 "contentFilter"
852 );
853 assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
854 assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
855 }
856
857 #[test]
858 fn token_usage_accumulates_cache_creation_prompt_tokens() {
859 let mut usage = TokenUsage::new(100, 10, 110)
860 .with_cached_prompt_tokens(80)
861 .with_cache_creation_prompt_tokens(15);
862 usage.add_assign(
863 &TokenUsage::new(50, 5, 55)
864 .with_cached_prompt_tokens(40)
865 .with_cache_creation_prompt_tokens(10),
866 );
867
868 assert_eq!(usage.prompt_tokens, 150);
869 assert_eq!(usage.cached_prompt_tokens, 120);
870 assert_eq!(usage.cache_creation_prompt_tokens, 25);
871 assert!(!usage.is_empty());
872
873 let creation_only = TokenUsage {
874 cache_creation_prompt_tokens: 1,
875 ..TokenUsage::default()
876 };
877 assert!(!creation_only.is_empty());
878 }
879
880 #[test]
881 fn service_tier_fields_default_when_absent_from_older_payloads() {
882 let hints: RuntimeHints = serde_json::from_value(serde_json::json!({
883 "trace_id": null,
884 "prompt_cache_key": null,
885 "auto_compact_token_limit": null
886 }))
887 .unwrap();
888 assert_eq!(hints.service_tier, None);
889 assert!(
890 serde_json::to_value(&hints)
891 .unwrap()
892 .get("service_tier")
893 .is_none()
894 );
895
896 let usage: TokenUsage = serde_json::from_value(serde_json::json!({
897 "prompt_tokens": 1,
898 "completion_tokens": 2,
899 "total_tokens": 3
900 }))
901 .unwrap();
902 assert_eq!(usage.service_tier, None);
903 }
904
905 #[test]
906 fn token_usage_keeps_the_most_recently_reported_service_tier() {
907 let mut usage = TokenUsage::new(10, 1, 11).with_service_tier(Some("priority".into()));
908 usage.add_assign(&TokenUsage::new(10, 1, 11));
909 assert_eq!(usage.service_tier.as_deref(), Some("priority"));
910 usage.add_assign(&TokenUsage::new(10, 1, 11).with_service_tier(Some("default".into())));
911 assert_eq!(usage.service_tier.as_deref(), Some("default"));
912 }
913
914 #[test]
915 fn inference_speed_policy_decision_serializes_runtime_metadata() {
916 let decision = SpeedPolicyDecision {
917 phase: SpeedPolicyPhase::Verification,
918 desired_reasoning: "high".to_string(),
919 applied_reasoning: Some("high".to_string()),
920 supported: true,
921 };
922 let hints = RuntimeHints {
923 speed_policy: Some(decision),
924 ..RuntimeHints::default()
925 };
926
927 let json = serde_json::to_value(hints).unwrap();
928 assert_eq!(
929 json.get("speed_policy")
930 .and_then(|value| value.get("phase"))
931 .and_then(serde_json::Value::as_str),
932 Some("verification")
933 );
934 assert_eq!(
935 json.get("speed_policy")
936 .and_then(|value| value.get("desiredReasoning"))
937 .and_then(serde_json::Value::as_str),
938 Some("high")
939 );
940 assert_eq!(
941 json.get("speed_policy")
942 .and_then(|value| value.get("appliedReasoning"))
943 .and_then(serde_json::Value::as_str),
944 Some("high")
945 );
946 }
947
948 #[test]
949 fn inference_reliability_policy_serializes_runtime_metadata() {
950 let hints = RuntimeHints {
951 reliability: Some(ReliabilityRequestPolicy::default()),
952 ..RuntimeHints::default()
953 };
954
955 let json = serde_json::to_value(hints).unwrap();
956 assert_eq!(
957 json.get("reliability")
958 .and_then(|value| value.get("providerRetryMaxAttempts"))
959 .and_then(serde_json::Value::as_u64),
960 Some(3)
961 );
962 assert_eq!(
963 json.get("reliability")
964 .and_then(|value| value.get("retryEmptyProviderBody"))
965 .and_then(serde_json::Value::as_bool),
966 Some(true)
967 );
968 }
969
970 #[test]
971 fn tool_search_config_serializes_provider_native_request() {
972 let config = ToolSearchConfig {
973 mode: ToolSearchMode::ProviderNative,
974 max_catalog_items: Some(200),
975 include_mcp: true,
976 include_skills: false,
977 fallback_to_explicit_tools: true,
978 provider_variant: ToolSearchProviderVariant::Bm25,
979 };
980
981 let value = serde_json::to_value(&config).unwrap();
982
983 assert_eq!(value["mode"], "provider_native");
984 assert_eq!(value["maxCatalogItems"], 200);
985 assert_eq!(value["includeMcp"], true);
986 assert_eq!(value["includeSkills"], false);
987 assert_eq!(value["providerVariant"], "bm25");
988 assert!(config.is_provider_native_requested());
989 }
990
991 #[test]
992 fn explicit_tool_search_config_preserves_current_default() {
993 let config = ToolSearchConfig::default();
994
995 assert_eq!(config.mode, ToolSearchMode::Explicit);
996 assert!(!config.is_provider_native_requested());
997 assert!(config.fallback_to_explicit_tools);
998 }
999
1000 #[test]
1001 fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
1002 let explicit = ToolSearchConfig::explicit();
1003 assert_eq!(
1004 explicit.resolve_effective_mode(true).unwrap(),
1005 EffectiveToolSearchMode::Explicit
1006 );
1007
1008 let auto = ToolSearchConfig {
1009 mode: ToolSearchMode::Auto,
1010 ..ToolSearchConfig::default()
1011 };
1012 assert_eq!(
1013 auto.resolve_effective_mode(true).unwrap(),
1014 EffectiveToolSearchMode::ProviderNative
1015 );
1016 assert_eq!(
1017 auto.resolve_effective_mode(false).unwrap(),
1018 EffectiveToolSearchMode::Explicit
1019 );
1020
1021 let native = ToolSearchConfig::provider_native();
1022 assert_eq!(
1023 native.resolve_effective_mode(true).unwrap(),
1024 EffectiveToolSearchMode::ProviderNative
1025 );
1026 assert_eq!(
1027 native.resolve_effective_mode(false).unwrap(),
1028 EffectiveToolSearchMode::Explicit
1029 );
1030
1031 let strict = ToolSearchConfig {
1032 fallback_to_explicit_tools: false,
1033 ..ToolSearchConfig::provider_native()
1034 };
1035 let error = strict.resolve_effective_mode(false).unwrap_err();
1036 assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
1037 assert!(error.to_string().contains("fallback_to_explicit_tools"));
1038 }
1039}