1use std::collections::BTreeMap;
2
3use serde::de::DeserializeOwned;
4use serde::{Deserialize, Serialize};
5use thiserror::Error;
6
7use crate::{AgentType, ModelPolicy, WorkflowUsage};
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum ModelRole {
12 Planner,
13 LeafReasoner,
14 Implementer,
15 Reviewer,
16 Teacher,
17 Student,
18 JsonExtractor,
19}
20
21impl From<AgentType> for ModelRole {
22 fn from(agent_type: AgentType) -> Self {
23 match agent_type {
24 AgentType::General | AgentType::Explore => Self::LeafReasoner,
25 AgentType::Plan => Self::Planner,
26 AgentType::Review | AgentType::Verifier => Self::Reviewer,
27 AgentType::Implementer => Self::Implementer,
28 }
29 }
30}
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
33pub struct ModelCapabilities {
34 #[serde(default)]
35 pub tool_calls: bool,
36 #[serde(default)]
37 pub json_mode: bool,
38 #[serde(default)]
39 pub prompt_cache: bool,
40 #[serde(default)]
41 pub large_context: bool,
42 #[serde(default)]
43 pub streaming: bool,
44}
45
46impl ModelCapabilities {
47 #[must_use]
48 pub fn satisfies(self, required: Self) -> bool {
49 (!required.tool_calls || self.tool_calls)
50 && (!required.json_mode || self.json_mode)
51 && (!required.prompt_cache || self.prompt_cache)
52 && (!required.large_context || self.large_context)
53 && (!required.streaming || self.streaming)
54 }
55}
56
57#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
58pub struct ProviderModel {
59 pub provider: String,
60 pub model: String,
61 #[serde(default)]
62 pub capabilities: ModelCapabilities,
63}
64
65#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
66pub struct ResolvedModel {
67 pub role: ModelRole,
68 pub provider: String,
69 pub model: String,
70 pub capabilities: ModelCapabilities,
71 pub source: ModelSelectionSource,
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
75#[serde(rename_all = "snake_case")]
76pub enum ModelSelectionSource {
77 Primary,
78 Fallback,
79 RoleDefault,
80}
81
82#[derive(Debug, Clone, Default)]
83pub struct ProviderRegistry {
84 models: BTreeMap<String, ProviderModel>,
85 role_policies: BTreeMap<ModelRole, ModelPolicy>,
86}
87
88impl ProviderRegistry {
89 pub fn new() -> Self {
90 Self::default()
91 }
92
93 pub fn with_model(mut self, model: ProviderModel) -> Self {
94 self.insert_model(model);
95 self
96 }
97
98 pub fn with_role_policy(mut self, role: ModelRole, policy: ModelPolicy) -> Self {
99 self.role_policies.insert(role, policy);
100 self
101 }
102
103 pub fn insert_model(&mut self, model: ProviderModel) {
104 self.models
105 .insert(model_key(&model.provider, &model.model), model);
106 }
107
108 pub fn resolve_role(
109 &self,
110 role: ModelRole,
111 policy: Option<&ModelPolicy>,
112 required: ModelCapabilities,
113 ) -> Result<ResolvedModel, ModelPolicyError> {
114 let policy = match policy {
115 Some(policy) => (policy, ModelSelectionSource::Primary),
116 None => (
117 self.role_policies
118 .get(&role)
119 .ok_or(ModelPolicyError::MissingPolicy { role })?,
120 ModelSelectionSource::RoleDefault,
121 ),
122 };
123 self.resolve_policy(role, policy.0, policy.1, required)
124 }
125
126 fn resolve_policy(
127 &self,
128 role: ModelRole,
129 policy: &ModelPolicy,
130 primary_source: ModelSelectionSource,
131 required: ModelCapabilities,
132 ) -> Result<ResolvedModel, ModelPolicyError> {
133 let candidates = model_candidates(policy)?;
134 let mut rejected = Vec::new();
135 for (index, candidate) in candidates.iter().enumerate() {
136 let source = if index == 0 {
137 primary_source
138 } else {
139 ModelSelectionSource::Fallback
140 };
141 let Some(model) = self
142 .models
143 .get(&model_key(&candidate.provider, &candidate.model))
144 else {
145 rejected.push(format!(
146 "{}/{}: unknown",
147 candidate.provider, candidate.model
148 ));
149 continue;
150 };
151 if model.capabilities.satisfies(required) {
152 return Ok(ResolvedModel {
153 role,
154 provider: model.provider.clone(),
155 model: model.model.clone(),
156 capabilities: model.capabilities,
157 source,
158 });
159 }
160 rejected.push(format!(
161 "{}/{}: missing required capabilities",
162 model.provider, model.model
163 ));
164 }
165 Err(ModelPolicyError::NoCapableModel { role, rejected })
166 }
167}
168
169#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
170pub struct CompletionRequest {
171 pub role: ModelRole,
172 pub prompt: String,
173 #[serde(default)]
174 pub require_json: bool,
175 #[serde(default)]
176 pub model_policy: ModelPolicy,
177}
178
179#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
180pub struct CompletionResponse {
181 pub text: String,
182 #[serde(default)]
183 pub usage: WorkflowUsage,
184}
185
186pub trait ModelProvider {
187 fn provider(&self) -> &str;
188 fn model(&self) -> &str;
189 fn capabilities(&self) -> ModelCapabilities;
190 fn complete(
191 &self,
192 request: &CompletionRequest,
193 ) -> Result<CompletionResponse, ModelProviderError>;
194}
195
196#[derive(Debug, Clone)]
197pub struct MockModelProvider {
198 provider: String,
199 model: String,
200 capabilities: ModelCapabilities,
201 response: CompletionResponse,
202}
203
204impl MockModelProvider {
205 pub fn new(
206 provider: impl Into<String>,
207 model: impl Into<String>,
208 capabilities: ModelCapabilities,
209 response: impl Into<String>,
210 ) -> Self {
211 Self {
212 provider: provider.into(),
213 model: model.into(),
214 capabilities,
215 response: CompletionResponse {
216 text: response.into(),
217 usage: WorkflowUsage::default(),
218 },
219 }
220 }
221}
222
223impl ModelProvider for MockModelProvider {
224 fn provider(&self) -> &str {
225 &self.provider
226 }
227
228 fn model(&self) -> &str {
229 &self.model
230 }
231
232 fn capabilities(&self) -> ModelCapabilities {
233 self.capabilities
234 }
235
236 fn complete(
237 &self,
238 _request: &CompletionRequest,
239 ) -> Result<CompletionResponse, ModelProviderError> {
240 Ok(self.response.clone())
241 }
242}
243
244#[derive(Debug, Clone, PartialEq, Eq, Error)]
245pub enum ModelPolicyError {
246 #[error("no model policy configured for role `{role:?}`")]
247 MissingPolicy { role: ModelRole },
248 #[error("model policy must include a model for role resolution")]
249 MissingModel,
250 #[error("fallback model `{model}` requires a provider when the primary policy has none")]
251 MissingFallbackProvider { model: String },
252 #[error("no configured model satisfies role `{role:?}` requirements: {rejected:?}")]
253 NoCapableModel {
254 role: ModelRole,
255 rejected: Vec<String>,
256 },
257}
258
259#[derive(Debug, Clone, PartialEq, Eq, Error)]
260pub enum ModelProviderError {
261 #[error("model provider `{provider}/{model}` failed: {reason}")]
262 Failed {
263 provider: String,
264 model: String,
265 reason: String,
266 },
267}
268
269#[derive(Debug, Clone, PartialEq, Eq, Error)]
270pub enum JsonRepairError {
271 #[error("json parse failed before and after one repair pass: {reason}")]
272 Parse { reason: String },
273}
274
275pub fn parse_json_with_repair<T: DeserializeOwned>(raw: &str) -> Result<T, JsonRepairError> {
276 match serde_json::from_str(raw) {
277 Ok(parsed) => Ok(parsed),
278 Err(first) => {
279 let repaired = repair_json_text_once(raw);
280 serde_json::from_str(&repaired).map_err(|second| JsonRepairError::Parse {
281 reason: format!("{first}; repair failed: {second}"),
282 })
283 }
284 }
285}
286
287pub fn repair_json_text_once(raw: &str) -> String {
288 let trimmed = raw.trim();
289 let without_fence = trimmed
290 .strip_prefix("```json")
291 .or_else(|| trimmed.strip_prefix("```"))
292 .and_then(|value| value.strip_suffix("```"))
293 .map(str::trim)
294 .unwrap_or(trimmed);
295 let object = slice_json_payload(without_fence, '{', '}');
296 let array = slice_json_payload(without_fence, '[', ']');
297 object.or(array).unwrap_or(without_fence).to_string()
298}
299
300#[derive(Debug, Clone, PartialEq, Eq)]
301struct ModelCandidate {
302 provider: String,
303 model: String,
304}
305
306fn model_candidates(policy: &ModelPolicy) -> Result<Vec<ModelCandidate>, ModelPolicyError> {
307 let mut candidates = Vec::new();
308 let Some(primary_model) = policy.model.as_ref() else {
309 return Err(ModelPolicyError::MissingModel);
310 };
311 candidates.push(candidate_from_model(
312 policy.provider.as_deref(),
313 primary_model,
314 )?);
315 for fallback in &policy.fallback_models {
316 candidates.push(candidate_from_model(policy.provider.as_deref(), fallback)?);
317 }
318 Ok(candidates)
319}
320
321fn candidate_from_model(
322 default_provider: Option<&str>,
323 model: &str,
324) -> Result<ModelCandidate, ModelPolicyError> {
325 if let Some((provider, model)) = model.split_once('/') {
326 return Ok(ModelCandidate {
327 provider: provider.to_string(),
328 model: model.to_string(),
329 });
330 }
331 let Some(provider) = default_provider else {
332 return Err(ModelPolicyError::MissingFallbackProvider {
333 model: model.to_string(),
334 });
335 };
336 Ok(ModelCandidate {
337 provider: provider.to_string(),
338 model: model.to_string(),
339 })
340}
341
342fn model_key(provider: &str, model: &str) -> String {
343 format!("{provider}/{model}")
344}
345
346fn slice_json_payload(raw: &str, open: char, close: char) -> Option<&str> {
347 let start = raw.find(open)?;
348 let end = raw.rfind(close)?;
349 (end >= start).then_some(&raw[start..=end])
350}
351
352#[cfg(test)]
353mod tests {
354 use super::*;
355
356 fn model(provider: &str, model: &str, capabilities: ModelCapabilities) -> ProviderModel {
357 ProviderModel {
358 provider: provider.to_string(),
359 model: model.to_string(),
360 capabilities,
361 }
362 }
363
364 #[test]
365 fn provider_capability_fallback() {
366 let registry = ProviderRegistry::new()
367 .with_model(model("mock", "plain", ModelCapabilities::default()))
368 .with_model(model(
369 "mock",
370 "json",
371 ModelCapabilities {
372 json_mode: true,
373 ..ModelCapabilities::default()
374 },
375 ));
376 let policy = ModelPolicy {
377 provider: Some("mock".to_string()),
378 model: Some("plain".to_string()),
379 fallback_models: vec!["json".to_string()],
380 };
381
382 let resolved = registry
383 .resolve_role(
384 ModelRole::JsonExtractor,
385 Some(&policy),
386 ModelCapabilities {
387 json_mode: true,
388 ..ModelCapabilities::default()
389 },
390 )
391 .expect("fallback json model should satisfy the role");
392
393 assert_eq!(resolved.model, "json");
394 assert_eq!(resolved.source, ModelSelectionSource::Fallback);
395 }
396
397 #[test]
398 fn role_default_policy_resolves_model() {
399 let registry = ProviderRegistry::new()
400 .with_model(model(
401 "mock",
402 "planner",
403 ModelCapabilities {
404 large_context: true,
405 ..ModelCapabilities::default()
406 },
407 ))
408 .with_role_policy(
409 ModelRole::Planner,
410 ModelPolicy {
411 provider: Some("mock".to_string()),
412 model: Some("planner".to_string()),
413 fallback_models: Vec::new(),
414 },
415 );
416
417 let resolved = registry
418 .resolve_role(
419 ModelRole::Planner,
420 None,
421 ModelCapabilities {
422 large_context: true,
423 ..ModelCapabilities::default()
424 },
425 )
426 .expect("role default should resolve");
427
428 assert_eq!(resolved.role, ModelRole::Planner);
429 assert_eq!(resolved.source, ModelSelectionSource::RoleDefault);
430 }
431
432 #[test]
433 fn agent_type_maps_to_model_role() {
434 assert_eq!(ModelRole::from(AgentType::Plan), ModelRole::Planner);
435 assert_eq!(
436 ModelRole::from(AgentType::Implementer),
437 ModelRole::Implementer
438 );
439 assert_eq!(ModelRole::from(AgentType::Verifier), ModelRole::Reviewer);
440 }
441
442 #[test]
443 fn json_repair_fallback() {
444 #[derive(Debug, Deserialize, PartialEq, Eq)]
445 struct Payload {
446 answer: String,
447 }
448
449 let parsed: Payload = parse_json_with_repair(
450 r#"Here is the JSON:
451```json
452{"answer":"ok"}
453```
454"#,
455 )
456 .expect("repair should extract fenced JSON");
457
458 assert_eq!(
459 parsed,
460 Payload {
461 answer: "ok".to_string()
462 }
463 );
464 }
465
466 #[test]
467 fn json_repair_fallback_fails_closed() {
468 let err = parse_json_with_repair::<serde_json::Value>("not json")
469 .expect_err("non-json text should fail closed");
470
471 assert!(matches!(err, JsonRepairError::Parse { .. }));
472 }
473
474 #[test]
475 fn mock_provider_returns_configured_response() {
476 let provider = MockModelProvider::new(
477 "mock",
478 "fast",
479 ModelCapabilities::default(),
480 "mock response",
481 );
482 let request = CompletionRequest {
483 role: ModelRole::LeafReasoner,
484 prompt: "say something".to_string(),
485 require_json: false,
486 model_policy: ModelPolicy::default(),
487 };
488
489 let response = provider.complete(&request).expect("mock should respond");
490
491 assert_eq!(provider.provider(), "mock");
492 assert_eq!(provider.model(), "fast");
493 assert_eq!(response.text, "mock response");
494 }
495}