Skip to main content

codewhale_workflow/
model_policy.rs

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
296    first_valid_json_payload(without_fence)
297        .unwrap_or(without_fence)
298        .to_string()
299}
300
301#[derive(Debug, Clone, PartialEq, Eq)]
302struct ModelCandidate {
303    provider: String,
304    model: String,
305}
306
307fn 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
322fn 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
343fn 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.
349const JSON_REPAIR_CANDIDATE_LIMIT: usize = 64;
350
351fn 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
372fn 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)]
410mod 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}