返回 CodeWhale
model_policy.rs
根目录 / crates / workflow / src / model_policy.rs
1 use std::collections::BTreeMap;
2
3 use serde::de::DeserializeOwned;
4 use serde::{Deserialize, Serialize};
5 use thiserror::Error;
6
7 use crate::{AgentType, ModelPolicy, WorkflowUsage};
8
9 #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
10 #[serde(rename_all = "snake_case")]
11 pub enum ModelRole {
12 Planner,
13 LeafReasoner,
14 Implementer,
15 Reviewer,
16 Teacher,
17 Student,
18 JsonExtractor,
19 }
20
21 impl 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)]
33 pub 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
46 impl 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)]
58 pub 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)]
66 pub 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")]
76 pub enum ModelSelectionSource {
77 Primary,
78 Fallback,
79 RoleDefault,
80 }
81
82 #[derive(Debug, Clone, Default)]
83 pub struct ProviderRegistry {
84 models: BTreeMap<String, ProviderModel>,
85 role_policies: BTreeMap<ModelRole, ModelPolicy>,
86 }
87
88 impl 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)]
170 pub 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)]
180 pub struct CompletionResponse {
181 pub text: String,
182 #[serde(default)]
183 pub usage: WorkflowUsage,
184 }
185
186 pub 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)]
197 pub struct MockModelProvider {
198 provider: String,
199 model: String,
200 capabilities: ModelCapabilities,
201 response: CompletionResponse,
202 }
203
204 impl 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
223 impl 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)]
245 pub 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)]
260 pub 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)]
270 pub enum JsonRepairError {
271 #[error("json parse failed before and after one repair pass: {reason}")]
272 Parse { reason: String },
273 }
274
275 pub 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
287 pub 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
296 first_valid_json_payload(without_fence)
297 .unwrap_or(without_fence)
298 .to_string()
299 }
300
301 #[derive(Debug, Clone, PartialEq, Eq)]
302 struct ModelCandidate {
303 provider: String,
304 model: String,
305 }
306
307 fn model_candidates(policy: &ModelPolicy) -> Result<Vec<ModelCandidate>, ModelPolicyError> {
308 let mut candidates = Vec::new();
309 let Some(primary_model) = policy.model.as_ref() else {
310 return Err(ModelPolicyError::MissingModel);
311 };
312 candidates.push(candidate_from_model(
313 policy.provider.as_deref(),
314 primary_model,
315 )?);
316 for fallback in &policy.fallback_models {
317 candidates.push(candidate_from_model(policy.provider.as_deref(), fallback)?);
318 }
319 Ok(candidates)
320 }
321
322 fn candidate_from_model(
323 default_provider: Option<&str>,
324 model: &str,
325 ) -> Result<ModelCandidate, ModelPolicyError> {
326 if let Some((provider, model)) = model.split_once('/') {
327 return Ok(ModelCandidate {
328 provider: provider.to_string(),
329 model: model.to_string(),
330 });
331 }
332 let Some(provider) = default_provider else {
333 return Err(ModelPolicyError::MissingFallbackProvider {
334 model: model.to_string(),
335 });
336 };
337 Ok(ModelCandidate {
338 provider: provider.to_string(),
339 model: model.to_string(),
340 })
341 }
342
343 fn model_key(provider: &str, model: &str) -> String {
344 format!("{provider}/{model}")
345 }
346
347 // Repair scans untrusted model text once per possible container start. Keep the
348 // fallback bounded when malformed output contains a long delimiter flood.
349 const JSON_REPAIR_CANDIDATE_LIMIT: usize = 64;
350
351 fn first_valid_json_payload(raw: &str) -> Option<&str> {
352 let mut attempted = 0;
353 for (start, open) in raw.char_indices() {
354 if !matches!(open, '{' | '[') {
355 continue;
356 }
357 attempted += 1;
358 if attempted > JSON_REPAIR_CANDIDATE_LIMIT {
359 break;
360 }
361
362 let Some(candidate) = balanced_json_payload(&raw[start..]) else {
363 continue;
364 };
365 if serde_json::from_str::<serde_json::Value>(candidate).is_ok() {
366 return Some(candidate);
367 }
368 }
369 None
370 }
371
372 fn balanced_json_payload(raw: &str) -> Option<&str> {
373 let mut stack = Vec::new();
374 let mut in_string = false;
375 let mut escaped = false;
376
377 for (offset, character) in raw.char_indices() {
378 if in_string {
379 if escaped {
380 escaped = false;
381 } else {
382 match character {
383 '\\' => escaped = true,
384 '"' => in_string = false,
385 _ => {}
386 }
387 }
388 continue;
389 }
390
391 match character {
392 '"' => in_string = true,
393 '{' | '[' => stack.push(character),
394 '}' | ']' => {
395 let expected_open = if character == '}' { '{' } else { '[' };
396 if stack.pop() != Some(expected_open) {
397 return None;
398 }
399 if stack.is_empty() {
400 return Some(&raw[..offset + character.len_utf8()]);
401 }
402 }
403 _ => {}
404 }
405 }
406 None
407 }
408
409 #[cfg(test)]
410 mod tests {
411 use super::*;
412
413 fn model(provider: &str, model: &str, capabilities: ModelCapabilities) -> ProviderModel {
414 ProviderModel {
415 provider: provider.to_string(),
416 model: model.to_string(),
417 capabilities,
418 }
419 }
420
421 #[test]
422 fn provider_capability_fallback() {
423 let registry = ProviderRegistry::new()
424 .with_model(model("mock", "plain", ModelCapabilities::default()))
425 .with_model(model(
426 "mock",
427 "json",
428 ModelCapabilities {
429 json_mode: true,
430 ..ModelCapabilities::default()
431 },
432 ));
433 let policy = ModelPolicy {
434 provider: Some("mock".to_string()),
435 model: Some("plain".to_string()),
436 fallback_models: vec!["json".to_string()],
437 };
438
439 let resolved = registry
440 .resolve_role(
441 ModelRole::JsonExtractor,
442 Some(&policy),
443 ModelCapabilities {
444 json_mode: true,
445 ..ModelCapabilities::default()
446 },
447 )
448 .expect("fallback json model should satisfy the role");
449
450 assert_eq!(resolved.model, "json");
451 assert_eq!(resolved.source, ModelSelectionSource::Fallback);
452 }
453
454 #[test]
455 fn role_default_policy_resolves_model() {
456 let registry = ProviderRegistry::new()
457 .with_model(model(
458 "mock",
459 "planner",
460 ModelCapabilities {
461 large_context: true,
462 ..ModelCapabilities::default()
463 },
464 ))
465 .with_role_policy(
466 ModelRole::Planner,
467 ModelPolicy {
468 provider: Some("mock".to_string()),
469 model: Some("planner".to_string()),
470 fallback_models: Vec::new(),
471 },
472 );
473
474 let resolved = registry
475 .resolve_role(
476 ModelRole::Planner,
477 None,
478 ModelCapabilities {
479 large_context: true,
480 ..ModelCapabilities::default()
481 },
482 )
483 .expect("role default should resolve");
484
485 assert_eq!(resolved.role, ModelRole::Planner);
486 assert_eq!(resolved.source, ModelSelectionSource::RoleDefault);
487 }
488
489 #[test]
490 fn agent_type_maps_to_model_role() {
491 assert_eq!(ModelRole::from(AgentType::Plan), ModelRole::Planner);
492 assert_eq!(
493 ModelRole::from(AgentType::Implementer),
494 ModelRole::Implementer
495 );
496 assert_eq!(ModelRole::from(AgentType::Verifier), ModelRole::Reviewer);
497 }
498
499 #[test]
500 fn json_repair_fallback() {
501 #[derive(Debug, Deserialize, PartialEq, Eq)]
502 struct Payload {
503 answer: String,
504 }
505
506 let parsed: Payload = parse_json_with_repair(
507 r#"Here is the JSON:
508 ```json
509 {"answer":"ok"}
510 ```
511 "#,
512 )
513 .expect("repair should extract fenced JSON");
514
515 assert_eq!(
516 parsed,
517 Payload {
518 answer: "ok".to_string()
519 }
520 );
521 }
522
523 #[test]
524 fn json_repair_fallback_fails_closed() {
525 let err = parse_json_with_repair::<serde_json::Value>("not json")
526 .expect_err("non-json text should fail closed");
527
528 assert!(matches!(err, JsonRepairError::Parse { .. }));
529 }
530
531 #[test]
532 fn mock_provider_returns_configured_response() {
533 let provider = MockModelProvider::new(
534 "mock",
535 "fast",
536 ModelCapabilities::default(),
537 "mock response",
538 );
539 let request = CompletionRequest {
540 role: ModelRole::LeafReasoner,
541 prompt: "say something".to_string(),
542 require_json: false,
543 model_policy: ModelPolicy::default(),
544 };
545
546 let response = provider.complete(&request).expect("mock should respond");
547
548 assert_eq!(provider.provider(), "mock");
549 assert_eq!(provider.model(), "fast");
550 assert_eq!(response.text, "mock response");
551 }
552
553 #[test]
554 fn repair_json_text_once_extracts_supported_payloads() {
555 // Plain JSON object
556 assert_eq!(
557 repair_json_text_once(r#"{"key": "value"}"#),
558 r#"{"key": "value"}"#
559 );
560
561 // Plain JSON array
562 assert_eq!(repair_json_text_once(r#"[1, 2, 3]"#), r#"[1, 2, 3]"#);
563
564 // Markdown fenced JSON object
565 assert_eq!(
566 repair_json_text_once("```json\n{\"key\": \"value\"}\n```"),
567 r#"{"key": "value"}"#
568 );
569
570 // Markdown fenced JSON array
571 assert_eq!(
572 repair_json_text_once("```json\n[1, 2, 3]\n```"),
573 r#"[1, 2, 3]"#
574 );
575
576 // Generic markdown fence
577 assert_eq!(
578 repair_json_text_once("```\n{\"key\": \"value\"}\n```"),
579 r#"{"key": "value"}"#
580 );
581
582 // JSON object embedded in text
583 assert_eq!(
584 repair_json_text_once("Here is the JSON:\n\n{\"key\": \"value\"}\n\nHope this helps!"),
585 r#"{"key": "value"}"#
586 );
587
588 // JSON array embedded in text
589 assert_eq!(
590 repair_json_text_once("Some text before [1, 2, 3] and some text after"),
591 r#"[1, 2, 3]"#
592 );
593
594 // Nested structures (object containing array)
595 assert_eq!(
596 repair_json_text_once(r#"{"key": [1, 2, 3]}"#),
597 r#"{"key": [1, 2, 3]}"#
598 );
599
600 // Nested structures (array containing object)
601 assert_eq!(
602 repair_json_text_once(r#"[{"key": "value"}]"#),
603 r#"[{"key": "value"}]"#
604 );
605
606 // Fenced JSON embedded in text
607 assert_eq!(
608 repair_json_text_once("Here is the JSON:\n```json\n{\"key\": \"value\"}\n```\nDone."),
609 r#"{"key": "value"}"#
610 );
611
612 // No valid JSON, fallback to trimmed text
613 assert_eq!(
614 repair_json_text_once("Just some plain text without json"),
615 "Just some plain text without json"
616 );
617
618 // Fenced plain text, falls back to stripped and trimmed text
619 assert_eq!(
620 repair_json_text_once("```json\nJust some plain text\n```"),
621 "Just some plain text"
622 );
623
624 // An unmatched opening bracket must not block a later valid object.
625 assert_eq!(repair_json_text_once("[note {\"ok\":1}"), r#"{"ok":1}"#);
626
627 // Delimiters inside strings must not terminate the outer object.
628 assert_eq!(
629 repair_json_text_once(r#"Prefix text {"data": "[nested]"} postfix text"#),
630 r#"{"data": "[nested]"}"#
631 );
632
633 // The earliest balanced candidate may still be invalid JSON. Keep scanning
634 // until a balanced candidate also parses successfully.
635 assert_eq!(
636 repair_json_text_once(r#"[not-json] then {"ok":true}"#),
637 r#"{"ok":true}"#
638 );
639
640 // Escaped quotes and backslashes stay inside the JSON string.
641 assert_eq!(
642 repair_json_text_once(
643 r#"prefix {"value":"a \"quoted\" [item]","path":"C:\\tmp"} suffix"#
644 ),
645 r#"{"value":"a \"quoted\" [item]","path":"C:\\tmp"}"#
646 );
647
648 // A mismatched outer delimiter must not consume a later valid payload.
649 assert_eq!(
650 repair_json_text_once(r#"[broken} then [1,{"ok":true}]"#),
651 r#"[1,{"ok":true}]"#
652 );
653 }
654 }
655
655 lines RUST