1use 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#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "lowercase")]
24pub enum SearchKind {
25 Markdown,
27 Schema,
29 Manifest,
31 Annotation,
33 Provenance,
35}
36
37impl SearchKind {
38 #[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#[derive(Clone, Debug, PartialEq, Eq)]
67pub struct SearchOptions {
68 pub limit: usize,
70 pub kind: Option<SearchKind>,
72 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
88pub struct SearchHit {
89 pub path: String,
91 pub kind: SearchKind,
93 #[serde(skip_serializing_if = "Option::is_none")]
95 pub heading: Option<String>,
96 #[serde(skip_serializing_if = "Option::is_none")]
98 pub line_start: Option<usize>,
99 #[serde(skip_serializing_if = "Option::is_none")]
101 pub line_end: Option<usize>,
102 pub score: f64,
104 pub text: String,
106}
107
108pub 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}