1use std::time::Duration;
37
38use crate::serde_json as crate_json;
39use crate::serde_json::Value as JsonValue;
40
41use crate::runtime::ask_pipeline::{extract_tokens as heuristic_extract_tokens, TokenSet};
42use crate::runtime::statement_frame::EffectiveScope;
43
44pub const DEFAULT_MAX_TOKENS: usize = 32;
47
48pub const DEFAULT_TIMEOUT_MS: u32 = 5_000;
51
52pub const NER_CAPABILITY: &str = "ai:ner:read";
54
55#[derive(Debug, Clone)]
60pub enum NerProvider {
61 OpenAiCompat { endpoint: String, model: String },
65 AnthropicNative { endpoint: String, model: String },
67 Stub(StubBehavior),
70}
71
72#[derive(Debug, Clone)]
74pub enum StubBehavior {
75 Empty,
77 Echo,
80 Canned(TokenSet),
83 SlowDuration(Duration),
86 RawJson(String),
91}
92
93#[derive(Debug, Clone, Copy, PartialEq, Eq)]
95pub enum HeuristicFallback {
96 UseHeuristic,
100 EmptyOnFail,
104 Propagate,
106}
107
108#[derive(Debug, Clone, PartialEq, Eq)]
111pub enum NerError {
112 NetworkTimeout,
114 ProviderRejected { status: u16, body_excerpt: String },
117 ResponseMalformed { reason: String },
119 ResponseExceedsTokenLimit { count: usize, max: usize },
121 SecretInResponse { pattern: String },
125 AuthDenied,
127}
128
129impl std::fmt::Display for NerError {
130 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
131 match self {
132 NerError::NetworkTimeout => write!(f, "ner: network timeout"),
133 NerError::ProviderRejected { status, .. } => {
134 write!(f, "ner: provider rejected (status={status})")
135 }
136 NerError::ResponseMalformed { reason } => {
137 write!(f, "ner: malformed response ({reason})")
138 }
139 NerError::ResponseExceedsTokenLimit { count, max } => {
140 write!(f, "ner: response exceeds token limit ({count} > {max})")
141 }
142 NerError::SecretInResponse { pattern } => {
143 write!(f, "ner: secret pattern in response ({pattern})")
144 }
145 NerError::AuthDenied => write!(f, "ner: auth denied (missing {NER_CAPABILITY})"),
146 }
147 }
148}
149
150impl std::error::Error for NerError {}
151
152pub trait AuthContext: std::fmt::Debug + Send + Sync {
159 fn has_capability(&self, capability: &str) -> bool;
160}
161
162#[derive(Debug, Clone, Default)]
166pub struct StubAuthContext {
167 capabilities: Vec<String>,
168}
169
170impl StubAuthContext {
171 pub fn new(caps: impl IntoIterator<Item = impl Into<String>>) -> Self {
172 Self {
173 capabilities: caps.into_iter().map(Into::into).collect(),
174 }
175 }
176
177 pub fn allow_all() -> Self {
178 Self::new([NER_CAPABILITY])
179 }
180
181 pub fn deny_all() -> Self {
182 Self::default()
183 }
184}
185
186impl AuthContext for StubAuthContext {
187 fn has_capability(&self, capability: &str) -> bool {
188 self.capabilities.iter().any(|c| c == capability)
189 }
190}
191
192#[derive(Debug, Clone)]
196pub struct LlmNer {
197 pub provider: NerProvider,
198 pub fallback: HeuristicFallback,
199 pub timeout_ms: u32,
200 pub max_tokens_returned: usize,
201}
202
203impl LlmNer {
204 pub fn new(provider: NerProvider, fallback: HeuristicFallback) -> Self {
206 Self {
207 provider,
208 fallback,
209 timeout_ms: DEFAULT_TIMEOUT_MS,
210 max_tokens_returned: DEFAULT_MAX_TOKENS,
211 }
212 }
213
214 pub async fn extract(
222 &self,
223 question: &str,
224 scope: &EffectiveScope,
225 auth: &dyn AuthContext,
226 ) -> Result<TokenSet, NerError> {
227 if !auth.has_capability(NER_CAPABILITY) {
231 return Err(NerError::AuthDenied);
232 }
233
234 let result = match &self.provider {
239 NerProvider::Stub(behavior) => self.run_stub(behavior, question),
240 NerProvider::OpenAiCompat { endpoint, model } => {
241 self.run_openai_compat(endpoint, model, question, scope)
242 .await
243 }
244 NerProvider::AnthropicNative { endpoint, model } => {
245 self.run_anthropic(endpoint, model, question, scope).await
246 }
247 };
248
249 match result {
250 Ok(tokens) => Ok(tokens),
251 Err(err) => self.handle_failure(err, question),
252 }
253 }
254
255 fn handle_failure(&self, err: NerError, question: &str) -> Result<TokenSet, NerError> {
257 if matches!(err, NerError::AuthDenied) {
259 return Err(err);
260 }
261 match self.fallback {
262 HeuristicFallback::UseHeuristic => Ok(heuristic_extract_tokens(question)),
263 HeuristicFallback::EmptyOnFail => Ok(TokenSet::default()),
264 HeuristicFallback::Propagate => Err(err),
265 }
266 }
267
268 fn run_stub(&self, behavior: &StubBehavior, question: &str) -> Result<TokenSet, NerError> {
270 match behavior {
271 StubBehavior::Empty => Ok(TokenSet::default()),
272 StubBehavior::Echo => {
273 let trimmed = question.trim().to_lowercase();
274 if trimmed.is_empty() {
275 Ok(TokenSet::default())
276 } else {
277 Ok(TokenSet {
278 keywords: vec![trimmed],
279 literals: vec![],
280 })
281 }
282 }
283 StubBehavior::Canned(tokens) => Ok(tokens.clone()),
284 StubBehavior::SlowDuration(d) => {
285 if d.as_millis() as u32 > self.timeout_ms {
289 Err(NerError::NetworkTimeout)
290 } else {
291 Ok(TokenSet::default())
292 }
293 }
294 StubBehavior::RawJson(raw) => parse_and_sanitize(raw, self.max_tokens_returned),
295 }
296 }
297
298 #[cfg(feature = "ai-ner-network")]
309 async fn run_openai_compat(
310 &self,
311 endpoint: &str,
312 model: &str,
313 question: &str,
314 scope: &EffectiveScope,
315 ) -> Result<TokenSet, NerError> {
316 let body = crate::json!({
317 "model": model,
318 "response_format": crate::json!({ "type": "json_object" }),
319 "messages": vec![
320 crate::json!({ "role": "system", "content": NER_SYSTEM_PROMPT }),
321 crate::json!({ "role": "user", "content": build_prompt(question, scope) }),
322 ],
323 });
324 let raw = http_post_json(endpoint, &body, self.timeout_ms).await?;
325 let payload = extract_openai_payload(&raw)?;
326 parse_and_sanitize(&payload, self.max_tokens_returned)
327 }
328
329 #[cfg(not(feature = "ai-ner-network"))]
330 async fn run_openai_compat(
331 &self,
332 _endpoint: &str,
333 _model: &str,
334 _question: &str,
335 _scope: &EffectiveScope,
336 ) -> Result<TokenSet, NerError> {
337 Err(NerError::NetworkTimeout)
340 }
341
342 #[cfg(feature = "ai-ner-network")]
343 async fn run_anthropic(
344 &self,
345 endpoint: &str,
346 model: &str,
347 question: &str,
348 scope: &EffectiveScope,
349 ) -> Result<TokenSet, NerError> {
350 let body = crate::json!({
351 "model": model,
352 "max_tokens": 1024,
353 "system": NER_SYSTEM_PROMPT,
354 "messages": vec![
355 crate::json!({ "role": "user", "content": build_prompt(question, scope) }),
356 ],
357 });
358 let raw = http_post_json(endpoint, &body, self.timeout_ms).await?;
359 let payload = extract_anthropic_payload(&raw)?;
360 parse_and_sanitize(&payload, self.max_tokens_returned)
361 }
362
363 #[cfg(not(feature = "ai-ner-network"))]
364 async fn run_anthropic(
365 &self,
366 _endpoint: &str,
367 _model: &str,
368 _question: &str,
369 _scope: &EffectiveScope,
370 ) -> Result<TokenSet, NerError> {
371 Err(NerError::NetworkTimeout)
372 }
373}
374
375const NER_SYSTEM_PROMPT: &str = "\
379You are an entity extraction service for a database query pipeline. \
380Read the user's question and return a JSON object with two fields: \
381'keywords' (array of lowercase content words, length >= 2) and \
382'literals' (array of identifier-shaped tokens kept in original case). \
383Return JSON only — no prose, no markdown.";
384
385#[allow(dead_code)] fn build_prompt(question: &str, scope: &EffectiveScope) -> String {
389 use crate::runtime::statement_frame::ReadFrame;
390 let visible: Vec<&str> = scope
391 .visible_collections()
392 .map(|set| set.iter().map(String::as_str).collect())
393 .unwrap_or_default();
394 format!(
395 "Question: {q}\nVisible collections: {v:?}\nReturn JSON only.",
396 q = question,
397 v = visible
398 )
399}
400
401#[cfg(feature = "ai-ner-network")]
402async fn http_post_json(
403 endpoint: &str,
404 body: &crate_json::Value,
405 timeout_ms: u32,
406) -> Result<String, NerError> {
407 let client = reqwest::Client::builder()
408 .timeout(Duration::from_millis(timeout_ms as u64))
409 .build()
410 .map_err(|e| NerError::ResponseMalformed {
411 reason: format!("client build: {e}"),
412 })?;
413 let resp = client
414 .post(endpoint)
415 .header("content-type", "application/json")
416 .body(body.to_string_compact())
417 .send()
418 .await
419 .map_err(|e| {
420 if e.is_timeout() {
421 NerError::NetworkTimeout
422 } else {
423 NerError::ResponseMalformed {
424 reason: format!("transport: {e}"),
425 }
426 }
427 })?;
428 let status = resp.status().as_u16();
429 let text = resp.text().await.map_err(|e| NerError::ResponseMalformed {
430 reason: format!("body read: {e}"),
431 })?;
432 if !(200..300).contains(&status) {
433 return Err(NerError::ProviderRejected {
434 status,
435 body_excerpt: scrub_excerpt(&text),
436 });
437 }
438 Ok(text)
439}
440
441#[cfg(feature = "ai-ner-network")]
442fn extract_openai_payload(raw: &str) -> Result<String, NerError> {
443 let v: JsonValue = crate_json::from_str(raw).map_err(|e| NerError::ResponseMalformed {
444 reason: format!("outer json: {e}"),
445 })?;
446 v["choices"]
447 .as_array()
448 .and_then(|choices| choices.first())
449 .and_then(|choice| choice["message"]["content"].as_str())
450 .map(str::to_owned)
451 .ok_or_else(|| NerError::ResponseMalformed {
452 reason: "missing choices[0].message.content".into(),
453 })
454}
455
456#[cfg(feature = "ai-ner-network")]
457fn extract_anthropic_payload(raw: &str) -> Result<String, NerError> {
458 let v: JsonValue = crate_json::from_str(raw).map_err(|e| NerError::ResponseMalformed {
459 reason: format!("outer json: {e}"),
460 })?;
461 v["content"]
462 .as_array()
463 .and_then(|content| content.first())
464 .and_then(|item| item["text"].as_str())
465 .map(str::to_owned)
466 .ok_or_else(|| NerError::ResponseMalformed {
467 reason: "missing content[0].text".into(),
468 })
469}
470
471#[allow(dead_code)] fn scrub_excerpt(s: &str) -> String {
473 let trimmed: String = s
474 .chars()
475 .take(256)
476 .filter(|c| !c.is_control() || *c == ' ')
477 .collect();
478 trimmed
479}
480
481fn parse_and_sanitize(raw: &str, max_tokens: usize) -> Result<TokenSet, NerError> {
487 let parsed: JsonValue = crate_json::from_str(raw).map_err(|e| NerError::ResponseMalformed {
488 reason: format!("json parse: {e}"),
489 })?;
490 let obj = parsed
491 .as_object()
492 .ok_or_else(|| NerError::ResponseMalformed {
493 reason: "expected JSON object at root".into(),
494 })?;
495
496 let keywords = collect_string_array(obj.get("keywords"), "keywords")?;
497 let literals = collect_string_array(obj.get("literals"), "literals")?;
498
499 let total = keywords.len() + literals.len();
500 if total > max_tokens {
501 return Err(NerError::ResponseExceedsTokenLimit {
502 count: total,
503 max: max_tokens,
504 });
505 }
506
507 for token in keywords.iter().chain(literals.iter()) {
508 validate_token(token)?;
509 }
510
511 Ok(TokenSet { keywords, literals })
512}
513
514fn collect_string_array(v: Option<&JsonValue>, field: &str) -> Result<Vec<String>, NerError> {
517 let arr = match v {
518 Some(JsonValue::Array(a)) => a,
519 Some(JsonValue::Null) | None => return Ok(Vec::new()),
520 Some(other) => {
521 return Err(NerError::ResponseMalformed {
522 reason: format!("{field}: expected array, got {}", json_kind(other)),
523 });
524 }
525 };
526 let mut out = Vec::with_capacity(arr.len());
527 for (i, item) in arr.iter().enumerate() {
528 match item {
529 JsonValue::String(s) => out.push(s.clone()),
530 other => {
531 return Err(NerError::ResponseMalformed {
532 reason: format!("{field}[{i}]: expected string, got {}", json_kind(other)),
533 });
534 }
535 }
536 }
537 Ok(out)
538}
539
540fn json_kind(v: &JsonValue) -> &'static str {
541 match v {
542 JsonValue::Null => "null",
543 JsonValue::Bool(_) => "bool",
544 JsonValue::Integer(_) => "number",
545 JsonValue::Number(_) => "number",
546 JsonValue::Decimal(_) => "number",
547 JsonValue::String(_) => "string",
548 JsonValue::Array(_) => "array",
549 JsonValue::Object(_) => "object",
550 }
551}
552
553fn validate_token(token: &str) -> Result<(), NerError> {
557 if let Some(pattern) = match_secret_pattern(token) {
558 return Err(NerError::SecretInResponse {
559 pattern: pattern.into(),
560 });
561 }
562 if token.is_empty() {
563 return Err(NerError::ResponseMalformed {
564 reason: "empty token".into(),
565 });
566 }
567 if token.len() > 256 {
568 return Err(NerError::ResponseMalformed {
569 reason: format!("token too long ({} bytes)", token.len()),
570 });
571 }
572 for (i, byte) in token.as_bytes().iter().enumerate() {
573 match byte {
574 0x00 => {
576 return Err(NerError::ResponseMalformed {
577 reason: format!("NUL byte at offset {i}"),
578 });
579 }
580 b'\n' | b'\r' => {
581 return Err(NerError::ResponseMalformed {
582 reason: format!("CR/LF at offset {i}"),
583 });
584 }
585 b'"' | b'\'' | b'`' => {
587 return Err(NerError::ResponseMalformed {
588 reason: format!("quote injection at offset {i}"),
589 });
590 }
591 b if *b < 0x20 && *b != b'\t' => {
593 return Err(NerError::ResponseMalformed {
594 reason: format!("control byte 0x{b:02x} at offset {i}"),
595 });
596 }
597 _ => {}
598 }
599 }
600 Ok(())
601}
602
603fn match_secret_pattern(token: &str) -> Option<&'static str> {
608 const PATTERNS: &[(&str, &str)] = &[
611 ("sk_", "sk_prefix"),
612 ("rs_", "rs_prefix"),
613 ("reddb_", "reddb_prefix"),
614 ("Bearer ", "bearer"),
615 ("bearer ", "bearer"),
616 ];
617 for (prefix, label) in PATTERNS {
618 if token.starts_with(prefix) {
619 return Some(label);
620 }
621 }
622 if looks_like_jwt(token) {
627 return Some("jwt");
628 }
629 if token.contains("://") && token.contains(':') && token.contains('@') {
631 if let Some(scheme_end) = token.find("://") {
632 let rest = &token[scheme_end + 3..];
633 if let Some(at) = rest.find('@') {
634 let userpass = &rest[..at];
635 if userpass.contains(':') {
636 return Some("conn_string_credentials");
637 }
638 }
639 }
640 }
641 None
642}
643
644fn looks_like_jwt(token: &str) -> bool {
645 let parts: Vec<&str> = token.split('.').collect();
646 if parts.len() != 3 {
647 return false;
648 }
649 parts.iter().all(|p| {
650 p.len() >= 4
651 && p.bytes()
652 .all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'-')
653 })
654}
655
656#[cfg(test)]
661mod tests {
662 use super::*;
663 use crate::runtime::ask_pipeline::TokenSet;
664
665 fn make_scope() -> EffectiveScope {
668 use crate::storage::transaction::snapshot::Snapshot;
672 use std::collections::HashSet;
673 EffectiveScope {
674 tenant: None,
675 identity: None,
676 snapshot: Snapshot {
677 xid: 0,
678 in_progress: HashSet::new(),
679 },
680 visible_collections: None,
681 }
682 }
683
684 fn allow() -> StubAuthContext {
685 StubAuthContext::allow_all()
686 }
687
688 fn deny() -> StubAuthContext {
689 StubAuthContext::deny_all()
690 }
691
692 #[tokio::test]
695 async fn stub_empty_returns_empty_token_set() {
696 let ner = LlmNer::new(
697 NerProvider::Stub(StubBehavior::Empty),
698 HeuristicFallback::Propagate,
699 );
700 let out = ner
701 .extract("anything", &make_scope(), &allow())
702 .await
703 .unwrap();
704 assert!(out.is_empty());
705 }
706
707 #[tokio::test]
708 async fn stub_echo_returns_lowercased_keyword() {
709 let ner = LlmNer::new(
710 NerProvider::Stub(StubBehavior::Echo),
711 HeuristicFallback::Propagate,
712 );
713 let out = ner
714 .extract(" Hello WORLD ", &make_scope(), &allow())
715 .await
716 .unwrap();
717 assert_eq!(out.keywords, vec!["hello world".to_string()]);
718 assert!(out.literals.is_empty());
719 }
720
721 #[tokio::test]
722 async fn stub_echo_empty_question_yields_empty_set() {
723 let ner = LlmNer::new(
724 NerProvider::Stub(StubBehavior::Echo),
725 HeuristicFallback::Propagate,
726 );
727 let out = ner.extract(" ", &make_scope(), &allow()).await.unwrap();
728 assert!(out.is_empty());
729 }
730
731 #[tokio::test]
732 async fn stub_canned_returns_provided_tokens() {
733 let canned = TokenSet {
734 keywords: vec!["passport".into()],
735 literals: vec!["FDD-1".into()],
736 };
737 let ner = LlmNer::new(
738 NerProvider::Stub(StubBehavior::Canned(canned.clone())),
739 HeuristicFallback::Propagate,
740 );
741 let out = ner.extract("q?", &make_scope(), &allow()).await.unwrap();
742 assert_eq!(out, canned);
743 }
744
745 #[tokio::test]
748 async fn slow_stub_within_budget_succeeds() {
749 let mut ner = LlmNer::new(
750 NerProvider::Stub(StubBehavior::SlowDuration(Duration::from_millis(10))),
751 HeuristicFallback::Propagate,
752 );
753 ner.timeout_ms = 100;
754 assert!(ner.extract("q?", &make_scope(), &allow()).await.is_ok());
755 }
756
757 #[tokio::test]
758 async fn slow_stub_over_budget_times_out_and_propagates() {
759 let mut ner = LlmNer::new(
760 NerProvider::Stub(StubBehavior::SlowDuration(Duration::from_millis(500))),
761 HeuristicFallback::Propagate,
762 );
763 ner.timeout_ms = 50;
764 let err = ner
765 .extract("q?", &make_scope(), &allow())
766 .await
767 .unwrap_err();
768 assert_eq!(err, NerError::NetworkTimeout);
769 }
770
771 #[tokio::test]
774 async fn malformed_not_json_is_rejected() {
775 let ner = LlmNer::new(
776 NerProvider::Stub(StubBehavior::RawJson("not-json".into())),
777 HeuristicFallback::Propagate,
778 );
779 let err = ner
780 .extract("q?", &make_scope(), &allow())
781 .await
782 .unwrap_err();
783 assert!(matches!(err, NerError::ResponseMalformed { .. }));
784 }
785
786 #[tokio::test]
787 async fn malformed_wrong_root_type_is_rejected() {
788 let ner = LlmNer::new(
789 NerProvider::Stub(StubBehavior::RawJson("[1,2,3]".into())),
790 HeuristicFallback::Propagate,
791 );
792 let err = ner
793 .extract("q?", &make_scope(), &allow())
794 .await
795 .unwrap_err();
796 assert!(matches!(err, NerError::ResponseMalformed { .. }));
797 }
798
799 #[tokio::test]
800 async fn malformed_keywords_not_array_is_rejected() {
801 let ner = LlmNer::new(
802 NerProvider::Stub(StubBehavior::RawJson(r#"{"keywords":"oops"}"#.into())),
803 HeuristicFallback::Propagate,
804 );
805 let err = ner
806 .extract("q?", &make_scope(), &allow())
807 .await
808 .unwrap_err();
809 assert!(matches!(err, NerError::ResponseMalformed { .. }));
810 }
811
812 fn adversarial_corpus() -> Vec<(&'static str, String)> {
818 let sk_prefix = format!("{}{}", "sk_", "live_DEADBEEFcafe");
820 let rs_prefix = format!("{}{}", "rs_", "test_TOKENtoken");
821 let reddb_prefix = format!("{}{}", "reddb_", "internal_secret_X");
822 let bearer = format!("{}{}", "Bearer ", "ABC.DEF.GHI");
823 let jwt = format!("{}.{}.{}", "abcd1234", "wxyz5678", "qrst9012");
824 let conn = "postgres://user:pwd@host:5432/db".to_string();
825
826 vec![
827 (
828 "crlf_in_keyword",
829 "{\"keywords\":[\"foo\\r\\nbar\"]}".into(),
830 ),
831 (
832 "nul_in_literal",
833 "{\"literals\":[\"foo\\u0000bar\"]}".into(),
834 ),
835 ("dquote_injection", "{\"keywords\":[\"foo\\\"bar\"]}".into()),
836 ("squote_injection", "{\"keywords\":[\"foo'bar\"]}".into()),
837 ("backtick_injection", "{\"keywords\":[\"foo`bar\"]}".into()),
838 (
839 "control_byte_low",
840 "{\"keywords\":[\"foo\\u0007bar\"]}".into(),
841 ),
842 ("sk_live", format!(r#"{{"keywords":["{sk_prefix}"]}}"#)),
843 ("rs_test", format!(r#"{{"keywords":["{rs_prefix}"]}}"#)),
844 (
845 "reddb_internal",
846 format!(r#"{{"literals":["{reddb_prefix}"]}}"#),
847 ),
848 ("bearer_token", format!(r#"{{"keywords":["{bearer}"]}}"#)),
849 ("jwt_shape", format!(r#"{{"literals":["{jwt}"]}}"#)),
850 ("conn_string", format!(r#"{{"keywords":["{conn}"]}}"#)),
851 ]
852 }
853
854 #[tokio::test]
855 async fn adversarial_corpus_is_fully_rejected() {
856 let corpus = adversarial_corpus();
857 assert!(corpus.len() >= 10, "corpus must be ≥10 payloads");
858 for (label, raw) in corpus {
859 let ner = LlmNer::new(
860 NerProvider::Stub(StubBehavior::RawJson(raw)),
861 HeuristicFallback::Propagate,
862 );
863 let err = ner
864 .extract("q?", &make_scope(), &allow())
865 .await
866 .expect_err(&format!("payload {label} should have been rejected"));
867 assert!(
868 matches!(
869 err,
870 NerError::ResponseMalformed { .. } | NerError::SecretInResponse { .. }
871 ),
872 "payload {label}: unexpected error variant {err:?}"
873 );
874 }
875 }
876
877 #[tokio::test]
878 async fn secret_in_response_reports_pattern_label() {
879 let raw = format!(r#"{{"keywords":["{}{}"]}}"#, "sk_", "live_zzzz");
880 let ner = LlmNer::new(
881 NerProvider::Stub(StubBehavior::RawJson(raw)),
882 HeuristicFallback::Propagate,
883 );
884 match ner
885 .extract("q?", &make_scope(), &allow())
886 .await
887 .unwrap_err()
888 {
889 NerError::SecretInResponse { pattern } => assert_eq!(pattern, "sk_prefix"),
890 other => panic!("expected SecretInResponse, got {other:?}"),
891 }
892 }
893
894 #[tokio::test]
897 async fn token_cap_excess_is_rejected() {
898 let kws: Vec<String> = (0..33).map(|i| format!("kw{i}")).collect();
900 let raw = crate_json::json!({ "keywords": kws }).to_string();
901 let ner = LlmNer::new(
902 NerProvider::Stub(StubBehavior::RawJson(raw)),
903 HeuristicFallback::Propagate,
904 );
905 let err = ner
906 .extract("q?", &make_scope(), &allow())
907 .await
908 .unwrap_err();
909 match err {
910 NerError::ResponseExceedsTokenLimit { count, max } => {
911 assert_eq!(count, 33);
912 assert_eq!(max, DEFAULT_MAX_TOKENS);
913 }
914 other => panic!("expected ResponseExceedsTokenLimit, got {other:?}"),
915 }
916 }
917
918 #[tokio::test]
919 async fn token_cap_at_limit_succeeds() {
920 let kws: Vec<String> = (0..DEFAULT_MAX_TOKENS).map(|i| format!("kw{i}")).collect();
921 let raw = crate_json::json!({ "keywords": kws }).to_string();
922 let ner = LlmNer::new(
923 NerProvider::Stub(StubBehavior::RawJson(raw)),
924 HeuristicFallback::Propagate,
925 );
926 let out = ner.extract("q?", &make_scope(), &allow()).await.unwrap();
927 assert_eq!(out.keywords.len(), DEFAULT_MAX_TOKENS);
928 }
929
930 #[tokio::test]
933 async fn auth_gate_denies_without_capability() {
934 let ner = LlmNer::new(
935 NerProvider::Stub(StubBehavior::Empty),
936 HeuristicFallback::UseHeuristic,
937 );
938 let err = ner.extract("q?", &make_scope(), &deny()).await.unwrap_err();
939 assert_eq!(err, NerError::AuthDenied);
940 }
941
942 #[tokio::test]
943 async fn auth_gate_denial_does_not_fall_back() {
944 let ner = LlmNer::new(
947 NerProvider::Stub(StubBehavior::Empty),
948 HeuristicFallback::UseHeuristic,
949 );
950 let err = ner
951 .extract("FDD-1", &make_scope(), &deny())
952 .await
953 .unwrap_err();
954 assert_eq!(err, NerError::AuthDenied);
955 }
956
957 #[tokio::test]
960 async fn fallback_use_heuristic_runs_extract_tokens() {
961 let ner = LlmNer::new(
963 NerProvider::Stub(StubBehavior::RawJson("not-json".into())),
964 HeuristicFallback::UseHeuristic,
965 );
966 let out = ner
967 .extract("show order 987654321 details", &make_scope(), &allow())
968 .await
969 .unwrap();
970 assert!(out.literals.iter().any(|l| l == "987654321"));
972 }
973
974 #[tokio::test]
975 async fn fallback_empty_on_fail_returns_empty() {
976 let ner = LlmNer::new(
977 NerProvider::Stub(StubBehavior::RawJson("not-json".into())),
978 HeuristicFallback::EmptyOnFail,
979 );
980 let out = ner
981 .extract("show order 987654321 details", &make_scope(), &allow())
982 .await
983 .unwrap();
984 assert!(out.is_empty());
985 }
986
987 #[tokio::test]
988 async fn fallback_propagate_returns_error() {
989 let ner = LlmNer::new(
990 NerProvider::Stub(StubBehavior::RawJson("not-json".into())),
991 HeuristicFallback::Propagate,
992 );
993 let err = ner
994 .extract("show order 987654321 details", &make_scope(), &allow())
995 .await
996 .unwrap_err();
997 assert!(matches!(err, NerError::ResponseMalformed { .. }));
998 }
999
1000 #[test]
1003 fn jwt_detector_matches_three_segments() {
1004 assert!(looks_like_jwt("abcd.efgh.ijkl"));
1005 assert!(!looks_like_jwt("abcd.efgh"));
1006 assert!(!looks_like_jwt("abc.def.ghi.jkl"));
1007 assert!(!looks_like_jwt("ab.cd.ef")); }
1009
1010 #[test]
1011 fn scrub_excerpt_drops_control_bytes() {
1012 let s = format!("ok\x07bad\nstill");
1013 let cleaned = scrub_excerpt(&s);
1014 assert!(!cleaned.contains('\x07'));
1015 assert!(!cleaned.contains('\n'));
1016 }
1017
1018 #[test]
1019 fn validate_token_accepts_normal_strings() {
1020 assert!(validate_token("passport").is_ok());
1021 assert!(validate_token("FDD-12313").is_ok());
1022 assert!(validate_token("foo_bar.baz").is_ok());
1023 }
1024}