Skip to main content

kimetsu_brain/
serving.rs

1//! Canonical brain-context retrieval and final MCP delivery policy. Evaluators
2//! supply the same request/configuration and measure this complete renderer.
3use crate::context::{
4    ContextBundle, ContextRequest,
5    delivery::{Delivery, compact_capsules, fit_json},
6    rerank_and_arbitrate,
7};
8use crate::embeddings::{Embedder, Reranker};
9use crate::project::BrainSession;
10use kimetsu_core::KimetsuResult;
11use serde_json::json;
12
13pub const RERANK_POOL: usize = 6;
14/// A score threshold, not a calibrated relevance probability.
15pub const RERANK_FLOOR: f32 = 0.30;
16pub const DEFAULT_BUDGET: u32 = 6000;
17pub const DEFAULT_CAP: usize = 3;
18/// ULIDs on the serving surface have exactly this length. Evaluation does not
19/// persist an exposure and therefore uses a deterministic placeholder.
20pub const EVAL_EXPOSURE_ID: &str = "00000000000000000000000000";
21
22/// Precompute once so a low-level fallback cannot disguise failed inference.
23struct QueryVector<'a> {
24    inner: &'a dyn Embedder,
25    vector: Vec<f32>,
26}
27impl Embedder for QueryVector<'_> {
28    fn embed(&self, _: &str) -> Result<Vec<f32>, crate::embeddings::EmbedderError> {
29        Ok(self.vector.clone())
30    }
31    fn model_id(&self) -> &str {
32        self.inner.model_id()
33    }
34    fn dim(&self) -> usize {
35        self.inner.dim()
36    }
37}
38struct CheckedScores<'a> {
39    model: &'a str,
40    scores: Vec<f32>,
41}
42impl Reranker for CheckedScores<'_> {
43    fn rerank(&self, _: &str, _: &[&str]) -> Result<Vec<f32>, crate::embeddings::EmbedderError> {
44        Ok(self.scores.clone())
45    }
46    fn model_id(&self) -> &str {
47        self.model
48    }
49}
50
51#[derive(Debug, Clone, Copy)]
52pub struct ServingPolicy {
53    pub budget: u32,
54    pub cap: usize,
55    pub pool: usize,
56    pub rerank_floor: f32,
57    pub explicit_fact_guard: bool,
58}
59impl Default for ServingPolicy {
60    fn default() -> Self {
61        Self {
62            budget: DEFAULT_BUDGET,
63            cap: DEFAULT_CAP,
64            pool: RERANK_POOL,
65            rerank_floor: RERANK_FLOOR,
66            explicit_fact_guard: false,
67        }
68    }
69}
70impl ServingPolicy {
71    pub fn from_config(config: &kimetsu_core::config::ProjectConfig) -> Self {
72        Self {
73            rerank_floor: config.broker.rerank_min_score,
74            explicit_fact_guard: config.broker.explicit_fact_guard,
75            ..Self::default()
76        }
77    }
78    pub fn prepare(&self, mut request: ContextRequest, reranking: bool) -> ContextRequest {
79        request.include_fact_evidence |= self.explicit_fact_guard;
80        request.defer_fact_budget =
81            self.explicit_fact_guard && crate::fact_query::parse(&request.query).is_some();
82        request.budget_tokens = if reranking {
83            self.budget.max(DEFAULT_BUDGET)
84        } else {
85            self.budget
86        };
87        request.max_capsules = if reranking || self.explicit_fact_guard {
88            self.cap.max(self.pool)
89        } else {
90            self.cap
91        };
92        request
93    }
94    pub fn arbitrate(
95        &self,
96        query: &str,
97        bundle: ContextBundle,
98        reranker: Option<&dyn Reranker>,
99        abstain: f32,
100    ) -> ContextBundle {
101        let mut bundle = rerank_and_arbitrate(
102            query,
103            bundle,
104            reranker,
105            abstain,
106            self.rerank_floor,
107            if self.explicit_fact_guard {
108                0
109            } else {
110                self.cap
111            },
112        );
113        if self.explicit_fact_guard {
114            crate::answerability::filter_bundle(query, &mut bundle);
115            if let Some(assessment) = crate::fact_query::evaluate(query, &bundle.capsules) {
116                bundle.known_fact_conflicts.extend(assessment.conflicting);
117                bundle.known_fact_conflicts.sort();
118                bundle.known_fact_conflicts.dedup();
119            }
120        }
121        if self.cap > 0 {
122            bundle.capsules.truncate(self.cap);
123        }
124        bundle.used_tokens = bundle.capsules.iter().map(|c| c.token_estimate).sum();
125        if bundle.capsules.is_empty() {
126            bundle.skipped = true;
127            bundle.evidence_coverage = 0.0;
128        }
129        bundle
130    }
131    pub fn render_for_query(
132        &self,
133        query: &str,
134        mut bundle: ContextBundle,
135        compress: bool,
136        exposure_id: &str,
137    ) -> Delivery {
138        if compress {
139            for capsule in &mut bundle.capsules {
140                capsule.summary = if self.explicit_fact_guard {
141                    crate::fact_query::compress_capsule(query, capsule, 3)
142                } else {
143                    crate::context::compress_for_render(&capsule.summary, 3)
144                };
145            }
146        }
147        if !self.explicit_fact_guard {
148            return self.render(bundle, false, exposure_id);
149        }
150        let mut known_conflicts = bundle.known_fact_conflicts.clone();
151        if let Some(assessment) = crate::fact_query::evaluate(query, &bundle.capsules) {
152            known_conflicts.extend(assessment.conflicting);
153        }
154        known_conflicts.sort();
155        known_conflicts.dedup();
156        let count = bundle.capsules.len();
157        fit_json(bundle.capsules.clone(), self.budget, |capsules| {
158            let mut payload = json!({
159                "ok":true,"skipped":capsules.is_empty(),"exposure_id":exposure_id,
160                "capsule_count":capsules.len(),"excluded_count":bundle.excluded.len()+count-capsules.len(),
161                "capsules":compact_capsules(capsules),"partial_evidence":bundle.evidence_coverage<1.0 || capsules.len()<count,
162            });
163            if let Some(mut assessment) = crate::fact_query::evaluate(query, capsules) {
164                crate::fact_query::preserve_conflicts(&mut assessment, &known_conflicts);
165                payload["partial_evidence"] =
166                    json!(assessment.status != "supported" || payload["partial_evidence"] == true);
167                payload["answerability"] = json!(assessment);
168            }
169            payload
170        })
171    }
172    pub fn render(&self, mut bundle: ContextBundle, compress: bool, exposure_id: &str) -> Delivery {
173        if compress {
174            for c in &mut bundle.capsules {
175                c.summary = crate::context::compress_for_render(&c.summary, 3)
176            }
177        }
178        let count = bundle.capsules.len();
179        fit_json(bundle.capsules.clone(), self.budget, |capsules| {
180            json!({
181                "ok":true,"skipped":capsules.is_empty(),"exposure_id":exposure_id,
182                "capsule_count":capsules.len(),"excluded_count":bundle.excluded.len()+count-capsules.len(),
183                "capsules":compact_capsules(capsules),"partial_evidence":bundle.evidence_coverage<1.0 || capsules.len()<count,
184            })
185        })
186    }
187    pub fn retrieve(
188        &self,
189        session: &BrainSession,
190        mut request: ContextRequest,
191        embedder: &dyn Embedder,
192        reranker: Option<&dyn Reranker>,
193        exposure_id: &str,
194    ) -> KimetsuResult<Delivery> {
195        session.resolve_request_floors(&mut request);
196        let abstain = request.abstain_evidence;
197        let query = request.query.clone();
198        let query_vector = if embedder.is_noop() {
199            None
200        } else {
201            let vector = embedder.embed(&query)?;
202            if vector.len() != embedder.dim()
203                || vector.is_empty()
204                || vector.iter().any(|x| !x.is_finite())
205                || !vector.iter().any(|x| *x != 0.0)
206            {
207                return Err(
208                    "embedder returned an invalid query vector; no semantic measurement".into(),
209                );
210            }
211            Some(QueryVector {
212                inner: embedder,
213                vector,
214            })
215        };
216        let checked_embedder = query_vector
217            .as_ref()
218            .map(|v| v as &dyn Embedder)
219            .unwrap_or(embedder);
220        let bundle = session.retrieve_context_with_injected_embedder(
221            self.prepare(request, reranker.is_some()),
222            checked_embedder,
223        )?;
224        let checked_scores = if let Some(rr) = reranker.filter(|_| !bundle.capsules.is_empty()) {
225            let docs: Vec<_> = bundle.capsules.iter().map(|c| c.summary.as_str()).collect();
226            let scores = rr.rerank(&query, &docs)?;
227            if scores.len() != docs.len()
228                || scores
229                    .iter()
230                    .any(|s| !s.is_finite() || !(0.0..=1.0).contains(s))
231            {
232                return Err(
233                    "reranker returned invalid scores; no cross-encoder measurement".into(),
234                );
235            }
236            Some(CheckedScores {
237                model: rr.model_id(),
238                scores,
239            })
240        } else {
241            None
242        };
243        let bundle = self.arbitrate(
244            &query,
245            bundle,
246            checked_scores.as_ref().map(|r| r as &dyn Reranker),
247            abstain,
248        );
249        Ok(self.render_for_query(
250            &query,
251            bundle,
252            session.config().broker.compress_capsules,
253            exposure_id,
254        ))
255    }
256}
257
258#[cfg(test)]
259mod tests {
260    use super::*;
261    use crate::{
262        context::ContextCapsule,
263        embeddings::{StubEmbedder, StubReranker},
264        project,
265    };
266    use kimetsu_core::memory::{MemoryKind, MemoryScope};
267    struct FailedEmbedder;
268    impl Embedder for FailedEmbedder {
269        fn embed(&self, _: &str) -> Result<Vec<f32>, crate::embeddings::EmbedderError> {
270            Err(crate::embeddings::EmbedderError::EmbedFailed(
271                "test failure".into(),
272            ))
273        }
274        fn model_id(&self) -> &str {
275            "failed"
276        }
277        fn dim(&self) -> usize {
278            2
279        }
280    }
281    struct MalformedReranker;
282    impl Reranker for MalformedReranker {
283        fn rerank(
284            &self,
285            _: &str,
286            _: &[&str],
287        ) -> Result<Vec<f32>, crate::embeddings::EmbedderError> {
288            Ok(vec![])
289        }
290        fn model_id(&self) -> &str {
291            "malformed"
292        }
293    }
294    #[test]
295    fn production_and_eval_use_same_final_budget_and_arbitration_with_injected_models() {
296        crate::user_brain::with_user_brain_disabled(|| {
297            let root = std::env::temp_dir().join(format!("kimetsu-serving-{}", ulid::Ulid::new()));
298            kimetsu_core::paths::git_init_boundary(&root);
299            project::init_project(&root, false).unwrap();
300            project::add_memory(
301                &root,
302                MemoryScope::Project,
303                MemoryKind::Fact,
304                "wal checkpoint protects sqlite commits",
305            )
306            .unwrap();
307            project::add_memory(
308                &root,
309                MemoryScope::Project,
310                MemoryKind::Fact,
311                "remote network bandwidth compression",
312            )
313            .unwrap();
314            let paths = kimetsu_core::paths::ProjectPaths::discover(&root).unwrap();
315            let mut config = project::load_config(&paths).unwrap();
316            config.broker.min_semantic_score = 0.73;
317            std::fs::write(&paths.project_toml, config.to_toml().unwrap()).unwrap();
318            let session = BrainSession::open_readonly(&root).unwrap();
319            let embedder = StubEmbedder::default();
320            let rr = StubReranker;
321            let request = ContextRequest {
322                query: "wal checkpoint".into(),
323                stage: "localization".into(),
324                ..Default::default()
325            };
326            assert!(
327                ServingPolicy::default()
328                    .retrieve(
329                        &session,
330                        request.clone(),
331                        &FailedEmbedder,
332                        None,
333                        EVAL_EXPOSURE_ID
334                    )
335                    .is_err(),
336                "failed inference cannot become successful semantic measurement"
337            );
338            assert!(
339                ServingPolicy::default()
340                    .retrieve(
341                        &session,
342                        request,
343                        &embedder,
344                        Some(&MalformedReranker),
345                        EVAL_EXPOSURE_ID
346                    )
347                    .is_err(),
348                "malformed CE cannot become successful reranker measurement"
349            );
350            for budget in [1, 600, 1200, 6000] {
351                let policy = ServingPolicy {
352                    budget,
353                    ..Default::default()
354                };
355                let req = ContextRequest {
356                    query: "wal checkpoint".into(),
357                    stage: "localization".into(),
358                    min_score: 0.15,
359                    ..Default::default()
360                };
361                let eval = policy
362                    .retrieve(
363                        &session,
364                        req.clone(),
365                        &embedder,
366                        Some(&rr),
367                        EVAL_EXPOSURE_ID,
368                    )
369                    .unwrap();
370                let mut resolved = req;
371                session.resolve_request_floors(&mut resolved);
372                let abstain = resolved.abstain_evidence;
373                let bundle = session
374                    .retrieve_context_with_injected_embedder(
375                        policy.prepare(resolved, true),
376                        &embedder,
377                    )
378                    .unwrap();
379                let production = policy.render(
380                    policy.arbitrate("wal checkpoint", bundle, Some(&rr), abstain),
381                    false,
382                    EVAL_EXPOSURE_ID,
383                );
384                let mut produced = production.payload.clone();
385                let mut measured = eval.payload.clone();
386                for payload in [&mut produced, &mut measured] {
387                    if let Some(caps) = payload["capsules"].as_array_mut() {
388                        for c in caps {
389                            c["id"] = json!(EVAL_EXPOSURE_ID);
390                        }
391                    }
392                }
393                assert_eq!(produced, measured);
394                assert_eq!(
395                    eval.payload["used_tokens"],
396                    json!(crate::context::delivery::serialized_output_tokens(
397                        &eval.payload
398                    ))
399                );
400                if budget == 1 {
401                    assert!(eval.capsules.is_empty());
402                    assert_eq!(eval.payload["error"], "budget_too_small");
403                }
404            }
405            let mut request = ContextRequest::default();
406            session.resolve_request_floors(&mut request);
407            assert_eq!(request.min_semantic_score, 0.73);
408            request.min_semantic_score_override = Some(0.0);
409            session.resolve_request_floors(&mut request);
410            assert_eq!(request.min_semantic_score, 0.0);
411            request.min_semantic_score_override = Some(-1.0);
412            session.resolve_request_floors(&mut request);
413            assert!(matches!(request.min_semantic_score, 0.35 | 0.0));
414        });
415    }
416    #[test]
417    fn compression_keeps_the_value_that_justified_admission() {
418        let capsule = ContextCapsule::wire_minimal(
419            "Background one. Background two. Background three. password = `test-value-only`".into(),
420            "memory".into(),
421            0.99,
422        );
423        let bundle = ContextBundle {
424            stage: "localization".into(),
425            budget_tokens: 6000,
426            used_tokens: 0,
427            capsules: vec![capsule],
428            excluded: vec![],
429            skipped: false,
430            top_score: 0.99,
431            top_abs_evidence: 0.99,
432            evidence_coverage: 1.0,
433            uncovered_terms: vec![],
434            chronological: false,
435            known_fact_conflicts: vec![],
436        };
437        let policy = ServingPolicy {
438            explicit_fact_guard: true,
439            ..Default::default()
440        };
441        let delivered =
442            policy.render_for_query("What is the password?", bundle, true, EVAL_EXPOSURE_ID);
443        assert!(delivered.payload.to_string().contains("test-value-only"));
444    }
445    #[test]
446    fn explicit_fact_guard_excludes_topic_match_before_output_cap() {
447        let capsules = [
448            "The listener binds port 6319.",
449            "password = `test-value-only`",
450        ]
451        .into_iter()
452        .map(|text| ContextCapsule::wire_minimal(text.into(), "memory".into(), 0.99))
453        .collect();
454        let bundle = ContextBundle {
455            stage: "localization".into(),
456            budget_tokens: 6000,
457            used_tokens: 0,
458            capsules,
459            excluded: vec![],
460            skipped: false,
461            top_score: 0.99,
462            top_abs_evidence: 0.99,
463            evidence_coverage: 1.0,
464            uncovered_terms: vec![],
465            chronological: false,
466            known_fact_conflicts: vec![],
467        };
468        let policy = ServingPolicy {
469            cap: 1,
470            explicit_fact_guard: true,
471            ..Default::default()
472        };
473        let selected = policy.arbitrate("What password is required?", bundle, None, 0.0);
474        assert_eq!(selected.capsules.len(), 1);
475        assert!(selected.capsules[0].summary.contains("test-value-only"));
476        assert_eq!(selected.excluded.len(), 1);
477    }
478
479    #[test]
480    fn reranker_floor_and_final_serialization_reject_candidates_before_measurement() {
481        let capsules = [("hit", "wal checkpoint"), ("noise", "remote network")]
482            .into_iter()
483            .map(|(id, text)| {
484                let mut c = ContextCapsule::wire_minimal(text.into(), "memory".into(), 0.8);
485                c.id = id.into();
486                c.expansion_handle = format!("memory:{id}");
487                c
488            })
489            .collect();
490        let bundle = ContextBundle {
491            stage: "localization".into(),
492            budget_tokens: 6000,
493            used_tokens: 0,
494            capsules,
495            excluded: vec![],
496            skipped: false,
497            top_score: 0.8,
498            top_abs_evidence: -1.0,
499            evidence_coverage: 1.0,
500            uncovered_terms: vec![],
501            chronological: false,
502            known_fact_conflicts: vec![],
503        };
504        let policy = ServingPolicy::default();
505        let selected = policy.arbitrate("wal checkpoint", bundle, Some(&StubReranker), 0.0);
506        assert_eq!(selected.capsules.len(), 1);
507        assert_eq!(selected.capsules[0].id, "hit");
508        let delivery = policy.render(selected, false, EVAL_EXPOSURE_ID);
509        assert_eq!(delivery.capsules.len(), 1);
510        assert!(!delivery.payload.to_string().contains("remote network"));
511    }
512}
513
514#[cfg(test)]
515mod conflict_carry_tests {
516    use super::*;
517    use crate::context::ContextCapsule;
518    struct RejectSecond;
519    impl Reranker for RejectSecond {
520        fn rerank(
521            &self,
522            _: &str,
523            _: &[&str],
524        ) -> Result<Vec<f32>, crate::embeddings::EmbedderError> {
525            Ok(vec![0.9, 0.0])
526        }
527        fn model_id(&self) -> &str {
528            "reject-second"
529        }
530    }
531    #[test]
532    fn guard_reserves_a_bounded_pool_before_the_final_cap() {
533        let policy = ServingPolicy {
534            cap: 1,
535            pool: 32,
536            explicit_fact_guard: true,
537            ..Default::default()
538        };
539        assert_eq!(
540            policy
541                .prepare(ContextRequest::default(), false)
542                .max_capsules,
543            32
544        );
545    }
546    #[test]
547    fn capsule_cap_does_not_turn_conflicting_evidence_into_support() {
548        let capsules = [("a", "7319"), ("b", "7320")]
549            .into_iter()
550            .map(|(id, value)| {
551                let text = format!("Orchid gateway port is {value}.");
552                let mut c = ContextCapsule::wire_minimal(text.clone(), "memory".into(), 0.99);
553                c.expansion_handle = format!("memory:{id}");
554                c.claim_revision = Some(format!("baseline:{id}"));
555                c.facts = crate::facts::extract(&text)
556                    .into_iter()
557                    .map(|claim| crate::fact_store::StoredFact {
558                        memory_id: id.into(),
559                        claim_revision: format!("baseline:{id}"),
560                        source_event_id: "source".into(),
561                        valid_from: None,
562                        valid_to: None,
563                        claim,
564                    })
565                    .collect();
566                c
567            })
568            .collect();
569        let bundle = ContextBundle {
570            stage: "localization".into(),
571            budget_tokens: 6000,
572            used_tokens: 0,
573            capsules,
574            excluded: vec![],
575            skipped: false,
576            top_score: 0.99,
577            top_abs_evidence: 0.99,
578            evidence_coverage: 1.0,
579            uncovered_terms: vec![],
580            chronological: false,
581            known_fact_conflicts: vec![],
582        };
583        let policy = ServingPolicy {
584            cap: 1,
585            budget: 6000,
586            explicit_fact_guard: true,
587            ..Default::default()
588        };
589        let q = "What is the Orchid gateway port?";
590        let eligible = policy.arbitrate(q, bundle.clone(), Some(&RejectSecond), 0.0);
591        let eligible_delivery = policy.render_for_query(q, eligible, true, EVAL_EXPOSURE_ID);
592        assert_eq!(
593            eligible_delivery.payload["answerability"]["status"],
594            "supported"
595        );
596        let chosen = policy.arbitrate(q, bundle, None, 0.0);
597        assert_eq!(chosen.capsules.len(), 1);
598        let delivery = policy.render_for_query(q, chosen, true, EVAL_EXPOSURE_ID);
599        assert_eq!(delivery.payload["answerability"]["status"], "conflicting");
600        assert_eq!(
601            delivery.payload["answerability"]["conflicting"],
602            serde_json::json!(["port"])
603        );
604    }
605}