Skip to main content

mcd_core/
search.rs

1//! In-memory BM25 search over package content and metadata.
2
3use std::collections::{HashMap, HashSet};
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7
8use crate::{
9    McdPackage,
10    annotations::load_manifest_annotations,
11    document::{DocumentBlock, McdDocument, SourceSpan},
12    markdown,
13    provenance::load_manifest_provenance,
14    schema::{TableColumnSchema, TableSchema},
15};
16
17const DEFAULT_LIMIT: usize = 10;
18const BM25_K1: f64 = 1.2;
19const BM25_B: f64 = 0.75;
20
21/// Search corpus kind.
22#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "lowercase")]
24pub enum SearchKind {
25    /// Markdown content blocks.
26    Markdown,
27    /// Table schema metadata.
28    Schema,
29    /// Manifest metadata.
30    Manifest,
31    /// Annotation metadata.
32    Annotation,
33    /// Provenance metadata.
34    Provenance,
35}
36
37impl SearchKind {
38    /// Parse a search kind name.
39    #[must_use]
40    pub fn parse(value: &str) -> Option<Self> {
41        match value {
42            "markdown" => Some(Self::Markdown),
43            "schema" => Some(Self::Schema),
44            "manifest" => Some(Self::Manifest),
45            "annotation" => Some(Self::Annotation),
46            "provenance" => Some(Self::Provenance),
47            _ => None,
48        }
49    }
50}
51
52impl std::fmt::Display for SearchKind {
53    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54        let value = match self {
55            Self::Markdown => "markdown",
56            Self::Schema => "schema",
57            Self::Manifest => "manifest",
58            Self::Annotation => "annotation",
59            Self::Provenance => "provenance",
60        };
61        f.write_str(value)
62    }
63}
64
65/// Search options.
66#[derive(Clone, Debug, PartialEq, Eq)]
67pub struct SearchOptions {
68    /// Maximum result count.
69    pub limit: usize,
70    /// Optional kind filter.
71    pub kind: Option<SearchKind>,
72    /// Optional package path filter.
73    pub page: Option<String>,
74}
75
76impl Default for SearchOptions {
77    fn default() -> Self {
78        Self {
79            limit: DEFAULT_LIMIT,
80            kind: None,
81            page: None,
82        }
83    }
84}
85
86/// Structured search result.
87#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
88pub struct SearchHit {
89    /// Package-internal source path.
90    pub path: String,
91    /// Hit kind.
92    pub kind: SearchKind,
93    /// Nearest Markdown heading or logical metadata title.
94    #[serde(skip_serializing_if = "Option::is_none")]
95    pub heading: Option<String>,
96    /// 1-based starting line when known.
97    #[serde(skip_serializing_if = "Option::is_none")]
98    pub line_start: Option<usize>,
99    /// 1-based ending line when known.
100    #[serde(skip_serializing_if = "Option::is_none")]
101    pub line_end: Option<usize>,
102    /// BM25 relevance score.
103    pub score: f64,
104    /// Indexed text snippet.
105    pub text: String,
106}
107
108/// Search a package with an in-memory BM25 index.
109pub fn search_package(
110    package: &McdPackage,
111    query: &str,
112    options: SearchOptions,
113) -> crate::Result<Vec<SearchHit>> {
114    let query_terms = unique_tokens(query);
115    if query_terms.is_empty() || options.limit == 0 {
116        return Ok(Vec::new());
117    }
118
119    let mut items = collect_corpus(package)?;
120    items.retain(|item| {
121        options.kind.is_none_or(|kind| item.kind == kind)
122            && options.page.as_deref().is_none_or(|page| item.path == page)
123    });
124    if items.is_empty() {
125        return Ok(Vec::new());
126    }
127
128    let document_count = items.len() as f64;
129    let tokenized = items
130        .iter()
131        .map(|item| tokenize(&item.search_text))
132        .collect::<Vec<_>>();
133    let average_len = tokenized
134        .iter()
135        .map(|tokens| tokens.len() as f64)
136        .sum::<f64>()
137        / document_count;
138
139    let mut document_frequency: HashMap<String, usize> = HashMap::new();
140    for tokens in &tokenized {
141        let seen = tokens.iter().map(String::as_str).collect::<HashSet<_>>();
142        for token in seen {
143            *document_frequency.entry(token.to_owned()).or_default() += 1;
144        }
145    }
146
147    let mut scored = items
148        .into_iter()
149        .zip(tokenized)
150        .filter_map(|(item, tokens)| {
151            let score = bm25_score(
152                &query_terms,
153                &tokens,
154                &document_frequency,
155                document_count,
156                average_len,
157            );
158            (score > 0.0).then_some((item, score))
159        })
160        .collect::<Vec<_>>();
161
162    scored.sort_by(|(left, left_score), (right, right_score)| {
163        right_score
164            .total_cmp(left_score)
165            .then_with(|| left.path.cmp(&right.path))
166            .then_with(|| left.kind.to_string().cmp(&right.kind.to_string()))
167            .then_with(|| left.line_start.cmp(&right.line_start))
168            .then_with(|| left.text.cmp(&right.text))
169    });
170
171    Ok(scored
172        .into_iter()
173        .take(options.limit)
174        .map(|(item, score)| SearchHit {
175            path: item.path,
176            kind: item.kind,
177            heading: item.heading,
178            line_start: item.line_start,
179            line_end: item.line_end,
180            score,
181            text: item.text,
182        })
183        .collect())
184}
185
186#[derive(Clone, Debug)]
187struct CorpusItem {
188    path: String,
189    kind: SearchKind,
190    heading: Option<String>,
191    line_start: Option<usize>,
192    line_end: Option<usize>,
193    text: String,
194    search_text: String,
195}
196
197fn collect_corpus(package: &McdPackage) -> crate::Result<Vec<CorpusItem>> {
198    let manifest = package.manifest()?;
199    let mut items = Vec::new();
200
201    for path in package
202        .entry_paths()
203        .into_iter()
204        .filter(|path| path.ends_with(".md"))
205    {
206        let markdown = package.read_to_string(path)?;
207        let document = markdown::parse_markdown(path, &markdown)?;
208        push_markdown_items(&mut items, &document);
209    }
210
211    push_manifest_items(&mut items, &manifest);
212    for table in &manifest.tables {
213        let schema = TableSchema::from_package(package, &table.schema)?;
214        push_schema_items(&mut items, &table.id, &table.schema, &schema);
215    }
216
217    let entry_document = McdDocument::from_package(package, &manifest)?;
218    let annotations = load_manifest_annotations(package, &manifest, &entry_document)?;
219    for (id, annotation) in annotations {
220        let value = serde_json::to_value(&annotation).unwrap_or(Value::Null);
221        items.push(CorpusItem {
222            path: manifest
223                .annotations
224                .iter()
225                .find(|entry| entry.id == id)
226                .map(|entry| entry.metadata.clone())
227                .unwrap_or_else(|| "manifest.json".to_owned()),
228            kind: SearchKind::Annotation,
229            heading: Some(id.clone()),
230            line_start: None,
231            line_end: None,
232            text: compact_join(json_strings(&value)),
233            search_text: compact_join(json_strings(&value)),
234        });
235    }
236
237    if let Some(provenance) = load_manifest_provenance(package, &manifest)?
238        && let Some(path) = &manifest.provenance
239    {
240        let value = serde_json::to_value(&provenance).unwrap_or(Value::Null);
241        let text = compact_join(json_strings(&value));
242        items.push(CorpusItem {
243            path: path.clone(),
244            kind: SearchKind::Provenance,
245            heading: Some("provenance".to_owned()),
246            line_start: None,
247            line_end: None,
248            text: text.clone(),
249            search_text: text,
250        });
251    }
252
253    Ok(items)
254}
255
256fn push_markdown_items(items: &mut Vec<CorpusItem>, document: &McdDocument) {
257    let mut current_heading: Option<String> = None;
258    for block in &document.blocks {
259        match block {
260            DocumentBlock::Heading { text, source, .. } => {
261                current_heading = Some(text.clone());
262                push_markdown_text(items, document, Some(text.clone()), *source, text.clone());
263            }
264            DocumentBlock::Paragraph { text, source, .. }
265            | DocumentBlock::List { text, source, .. }
266            | DocumentBlock::Quote { text, source, .. }
267            | DocumentBlock::MathBlock { text, source, .. } => {
268                push_markdown_text(
269                    items,
270                    document,
271                    current_heading.clone(),
272                    *source,
273                    text.clone(),
274                );
275            }
276            DocumentBlock::CodeBlock {
277                text,
278                source,
279                language,
280                ..
281            } => {
282                let display = language
283                    .as_deref()
284                    .map(|language| format!("{language}\n{text}"))
285                    .unwrap_or_else(|| text.clone());
286                push_markdown_text(items, document, current_heading.clone(), *source, display);
287            }
288            DocumentBlock::TableRef {
289                placement, source, ..
290            } => {
291                let text = [
292                    placement.ref_id.as_deref(),
293                    Some(placement.table.as_str()),
294                    placement.view.as_deref(),
295                    placement.caption.as_deref(),
296                ]
297                .into_iter()
298                .flatten()
299                .collect::<Vec<_>>()
300                .join(" ");
301                push_markdown_text(items, document, current_heading.clone(), *source, text);
302            }
303            DocumentBlock::ImageRef {
304                placement, source, ..
305            } => {
306                let text = [
307                    placement.ref_id.as_deref(),
308                    placement.asset.as_deref(),
309                    placement.image.as_deref(),
310                    placement.alt.as_deref(),
311                    placement.caption.as_deref(),
312                ]
313                .into_iter()
314                .flatten()
315                .collect::<Vec<_>>()
316                .join(" ");
317                push_markdown_text(items, document, current_heading.clone(), *source, text);
318            }
319        }
320    }
321}
322
323fn push_markdown_text(
324    items: &mut Vec<CorpusItem>,
325    document: &McdDocument,
326    heading: Option<String>,
327    source: Option<SourceSpan>,
328    text: String,
329) {
330    if text.trim().is_empty() {
331        return;
332    }
333    items.push(CorpusItem {
334        path: document.source_path.clone(),
335        kind: SearchKind::Markdown,
336        heading,
337        line_start: source.map(|source| source.start_line),
338        line_end: source.map(|source| source.end_line),
339        search_text: text.clone(),
340        text,
341    });
342}
343
344fn push_manifest_items(items: &mut Vec<CorpusItem>, manifest: &crate::Manifest) {
345    let mut parts = vec![
346        manifest.format.clone(),
347        manifest.version.clone(),
348        manifest.entrypoint.clone(),
349    ];
350    if let Some(title) = &manifest.title {
351        parts.push(title.clone());
352    }
353    parts.extend(manifest.tables.iter().flat_map(|table| {
354        [
355            table.id.clone(),
356            table.data.clone(),
357            table.schema.clone(),
358            table.views.keys().cloned().collect::<Vec<_>>().join(" "),
359            table.views.values().cloned().collect::<Vec<_>>().join(" "),
360        ]
361    }));
362    parts.extend(
363        manifest
364            .images
365            .iter()
366            .flat_map(|image| [image.id.clone(), image.metadata.clone()]),
367    );
368    parts.extend(manifest.external_data.iter().flat_map(|item| {
369        [
370            item.id.clone(),
371            item.uri.clone(),
372            item.media_type.clone(),
373            item.description.clone().unwrap_or_default(),
374        ]
375    }));
376    let text = compact_join(parts);
377    if !text.is_empty() {
378        items.push(CorpusItem {
379            path: "manifest.json".to_owned(),
380            kind: SearchKind::Manifest,
381            heading: manifest
382                .title
383                .clone()
384                .or_else(|| Some("manifest".to_owned())),
385            line_start: None,
386            line_end: None,
387            text: text.clone(),
388            search_text: text,
389        });
390    }
391}
392
393fn push_schema_items(
394    items: &mut Vec<CorpusItem>,
395    table_id: &str,
396    path: &str,
397    schema: &TableSchema,
398) {
399    let primary_key = if schema.primary_key.is_empty() {
400        String::new()
401    } else {
402        format!("primary key {}", schema.primary_key.join(" "))
403    };
404    let table_text = compact_join([
405        table_id.to_owned(),
406        schema.id.clone(),
407        primary_key,
408        schema
409            .foreign_keys
410            .iter()
411            .map(|key| {
412                format!(
413                    "foreign key {} references {} {}",
414                    key.columns.join(" "),
415                    key.references.table,
416                    key.references.columns.join(" ")
417                )
418            })
419            .collect::<Vec<_>>()
420            .join(" "),
421    ]);
422    if !table_text.is_empty() {
423        items.push(CorpusItem {
424            path: path.to_owned(),
425            kind: SearchKind::Schema,
426            heading: Some(table_id.to_owned()),
427            line_start: None,
428            line_end: None,
429            text: table_text.clone(),
430            search_text: table_text,
431        });
432    }
433
434    for column in &schema.columns {
435        let text = column_text(table_id, column);
436        items.push(CorpusItem {
437            path: path.to_owned(),
438            kind: SearchKind::Schema,
439            heading: Some(format!("{table_id}.{}", column.name)),
440            line_start: None,
441            line_end: None,
442            text: text.clone(),
443            search_text: text,
444        });
445    }
446}
447
448fn column_text(table_id: &str, column: &TableColumnSchema) -> String {
449    let mut parts = vec![
450        table_id.to_owned(),
451        column.name.clone(),
452        column.value_type.to_string(),
453    ];
454    if let Some(label) = &column.label {
455        parts.push(label.clone());
456    }
457    if let Some(unit) = &column.unit {
458        if let Some(code) = &unit.code {
459            parts.push(code.clone());
460        }
461        if let Some(label) = &unit.label {
462            parts.push(label.clone());
463        }
464    }
465    parts.extend(column.enum_values.clone());
466    compact_join(parts)
467}
468
469fn bm25_score(
470    query_terms: &[String],
471    tokens: &[String],
472    document_frequency: &HashMap<String, usize>,
473    document_count: f64,
474    average_len: f64,
475) -> f64 {
476    if tokens.is_empty() || average_len == 0.0 {
477        return 0.0;
478    }
479    let mut term_frequency: HashMap<&str, usize> = HashMap::new();
480    for token in tokens {
481        *term_frequency.entry(token).or_default() += 1;
482    }
483    let document_len = tokens.len() as f64;
484    query_terms
485        .iter()
486        .filter_map(|term| {
487            let frequency = *term_frequency.get(term.as_str())? as f64;
488            let document_frequency = *document_frequency.get(term.as_str()).unwrap_or(&0) as f64;
489            let idf = (1.0
490                + (document_count - document_frequency + 0.5) / (document_frequency + 0.5))
491                .ln();
492            let denominator =
493                frequency + BM25_K1 * (1.0 - BM25_B + BM25_B * document_len / average_len);
494            Some(idf * frequency * (BM25_K1 + 1.0) / denominator)
495        })
496        .sum()
497}
498
499fn unique_tokens(text: &str) -> Vec<String> {
500    let mut seen = HashSet::new();
501    tokenize(text)
502        .into_iter()
503        .filter(|token| seen.insert(token.clone()))
504        .collect()
505}
506
507fn tokenize(text: &str) -> Vec<String> {
508    let mut tokens = Vec::new();
509    let mut current = String::new();
510    for character in text.chars() {
511        if character.is_ascii_alphanumeric() || character == '_' {
512            current.push(character.to_ascii_lowercase());
513        } else {
514            push_token_parts(&mut tokens, &mut current);
515        }
516    }
517    push_token_parts(&mut tokens, &mut current);
518    tokens
519}
520
521fn push_token_parts(tokens: &mut Vec<String>, current: &mut String) {
522    if current.is_empty() {
523        return;
524    }
525    tokens.push(current.clone());
526    if current.contains('_') {
527        tokens.extend(
528            current
529                .split('_')
530                .filter(|part| !part.is_empty())
531                .map(ToOwned::to_owned),
532        );
533    }
534    current.clear();
535}
536
537fn json_strings(value: &Value) -> Vec<String> {
538    let mut strings = Vec::new();
539    collect_json_strings(value, &mut strings);
540    strings
541}
542
543fn collect_json_strings(value: &Value, strings: &mut Vec<String>) {
544    match value {
545        Value::String(value) if !value.trim().is_empty() => strings.push(value.clone()),
546        Value::Array(values) => {
547            for value in values {
548                collect_json_strings(value, strings);
549            }
550        }
551        Value::Object(values) => {
552            for (key, value) in values {
553                strings.push(key.clone());
554                collect_json_strings(value, strings);
555            }
556        }
557        Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => {}
558    }
559}
560
561fn compact_join(parts: impl IntoIterator<Item = String>) -> String {
562    parts
563        .into_iter()
564        .filter(|part| !part.trim().is_empty())
565        .collect::<Vec<_>>()
566        .join(" ")
567}
568
569impl std::fmt::Display for crate::schema::ColumnType {
570    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
571        let value = match self {
572            Self::String => "string",
573            Self::Integer => "integer",
574            Self::Decimal => "decimal",
575            Self::Boolean => "boolean",
576            Self::Date => "date",
577            Self::Datetime => "datetime",
578            Self::Time => "time",
579            Self::Enum => "enum",
580        };
581        f.write_str(value)
582    }
583}
584
585#[cfg(test)]
586mod tests {
587    use super::*;
588    use std::io::{Cursor, Write};
589    use zip::{CompressionMethod, ZipWriter, write::SimpleFileOptions};
590
591    #[test]
592    fn searches_markdown_and_schema_without_rows() {
593        let package = McdPackage::from_bytes(&zip_bytes(&[
594            ("mimetype", crate::package::MCD_MIMETYPE),
595            (
596                "manifest.json",
597                r#"{
598                    "format":"MCD",
599                    "version":"0.1",
600                    "profile":"MCD-Core",
601                    "entrypoint":"content/main.md",
602                    "title":"Thermal Dossier",
603                    "tables":[{"id":"powertrain","data":"tables/powertrain.csv","schema":"tables/powertrain.schema.json"}]
604                }"#,
605            ),
606            (
607                "content/main.md",
608                "# Powertrain calibration specifications\n\nThe `thermal_limit_deg_c` field constrains coolant flow for V50D.\n",
609            ),
610            ("tables/powertrain.csv", "calibration_id,engine_family\nCAL-1,V50D\n"),
611            (
612                "tables/powertrain.schema.json",
613                r#"{"id":"powertrain","columns":[{"name":"calibration_id","type":"string","label":"Calibration ID"},{"name":"thermal_limit_deg_c","type":"decimal","label":"Thermal Limit","unit":{"code":"deg_C","label":"deg C"}}]}"#,
614            ),
615        ]))
616        .expect("package opens");
617
618        let hits = search_package(
619            &package,
620            "thermal_limit_deg_c coolant V50D",
621            SearchOptions {
622                limit: 5,
623                kind: None,
624                page: None,
625            },
626        )
627        .expect("search succeeds");
628
629        assert!(hits.iter().any(|hit| {
630            hit.kind == SearchKind::Markdown
631                && hit.path == "content/main.md"
632                && hit.line_start == Some(3)
633        }));
634        assert!(hits.iter().any(|hit| {
635            hit.kind == SearchKind::Schema
636                && hit.path == "tables/powertrain.schema.json"
637                && hit.text.contains("thermal_limit_deg_c")
638        }));
639        assert!(!hits.iter().any(|hit| hit.text.contains("CAL-1")));
640    }
641
642    #[test]
643    fn filters_kind_and_page() {
644        let package = McdPackage::from_markdown("# Title\n\nA coolant paragraph.\n");
645        let hits = package
646            .search(
647                "coolant",
648                SearchOptions {
649                    limit: 10,
650                    kind: Some(SearchKind::Markdown),
651                    page: Some("content/main.md".to_owned()),
652                },
653            )
654            .expect("search succeeds");
655
656        assert_eq!(hits.len(), 1);
657        assert_eq!(hits[0].line_start, Some(3));
658    }
659
660    fn zip_bytes(entries: &[(&str, &str)]) -> Vec<u8> {
661        let cursor = Cursor::new(Vec::new());
662        let mut writer = ZipWriter::new(cursor);
663        let options = SimpleFileOptions::default().compression_method(CompressionMethod::Stored);
664
665        for (path, content) in entries {
666            writer.start_file(*path, options).expect("start file");
667            writer.write_all(content.as_bytes()).expect("write file");
668        }
669
670        writer.finish().expect("finish zip").into_inner()
671    }
672}