1use crate::ai::AiConfig;
7use crate::model::{Literal, NamedNode, Triple};
8use anyhow::Result;
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11
12pub struct RelationExtractor {
14 config: ExtractionConfig,
16
17 ner_model: Box<dyn NamedEntityRecognizer>,
19
20 relation_model: Box<dyn RelationClassifier>,
22
23 entity_linker: Box<dyn EntityLinker>,
25
26 confidence_threshold: f32,
28}
29
30#[derive(Debug, Clone, Serialize, Deserialize)]
32pub struct ExtractionConfig {
33 pub enable_ner: bool,
35
36 pub enable_relation_classification: bool,
38
39 pub enable_entity_linking: bool,
41
42 pub confidence_threshold: f32,
44
45 pub max_sentence_length: usize,
47
48 pub language_model: String,
50
51 pub enable_coreference: bool,
53
54 pub supported_languages: Vec<String>,
56}
57
58impl Default for ExtractionConfig {
59 fn default() -> Self {
60 Self {
61 enable_ner: true,
62 enable_relation_classification: true,
63 enable_entity_linking: true,
64 confidence_threshold: 0.7,
65 max_sentence_length: 512,
66 language_model: "bert-base-uncased".to_string(),
67 enable_coreference: true,
68 supported_languages: vec!["en".to_string()],
69 }
70 }
71}
72
73#[derive(Debug, Clone, Serialize, Deserialize)]
75pub struct ExtractedRelation {
76 pub subject: ExtractedEntity,
78
79 pub predicate: String,
81
82 pub object: ExtractedEntity,
84
85 pub confidence: f32,
87
88 pub source_span: TextSpan,
90
91 pub context: String,
93
94 pub metadata: HashMap<String, String>,
96}
97
98#[derive(Debug, Clone, Serialize, Deserialize)]
100pub struct ExtractedEntity {
101 pub text: String,
103
104 pub entity_type: EntityType,
106
107 pub kb_id: Option<String>,
109
110 pub confidence: f32,
112
113 pub span: TextSpan,
115}
116
117#[derive(Debug, Clone, Serialize, Deserialize)]
119pub enum EntityType {
120 Person,
121 Organization,
122 Location,
123 Date,
124 Time,
125 Money,
126 Percent,
127 Product,
128 Event,
129 Concept,
130 Other(String),
131}
132
133#[derive(Debug, Clone, Serialize, Deserialize)]
135pub struct TextSpan {
136 pub start: usize,
138
139 pub end: usize,
141
142 pub text: String,
144}
145
146pub trait NamedEntityRecognizer: Send + Sync {
148 fn extract_entities(&self, text: &str) -> Result<Vec<ExtractedEntity>>;
150
151 fn supported_types(&self) -> Vec<EntityType>;
153}
154
155pub trait RelationClassifier: Send + Sync {
157 fn classify_relation(
159 &self,
160 text: &str,
161 subject: &ExtractedEntity,
162 object: &ExtractedEntity,
163 ) -> Result<Option<(String, f32)>>;
164
165 fn supported_relations(&self) -> Vec<String>;
167}
168
169pub trait EntityLinker: Send + Sync {
171 fn link_entity(&self, entity: &ExtractedEntity, context: &str) -> Result<Option<String>>;
173
174 fn kb_info(&self) -> KnowledgeBaseInfo;
176}
177
178#[derive(Debug, Clone, Serialize, Deserialize)]
180pub struct KnowledgeBaseInfo {
181 pub name: String,
183
184 pub base_uri: String,
186
187 pub version: String,
189
190 pub entity_count: usize,
192}
193
194impl RelationExtractor {
195 pub fn new(_config: &AiConfig) -> Result<Self> {
213 Ok(Self {
214 config: ExtractionConfig::default(),
215 ner_model: Box::new(HeuristicNer::new()),
216 relation_model: Box::new(HeuristicRelationClassifier::new()),
217 entity_linker: Box::new(LocalEntityLinker::new()),
218 confidence_threshold: 0.7,
219 })
220 }
221
222 pub fn with_backends(
228 extraction_config: ExtractionConfig,
229 ner_model: Box<dyn NamedEntityRecognizer>,
230 relation_model: Box<dyn RelationClassifier>,
231 entity_linker: Box<dyn EntityLinker>,
232 ) -> Self {
233 let confidence_threshold = extraction_config.confidence_threshold;
234 Self {
235 config: extraction_config,
236 ner_model,
237 relation_model,
238 entity_linker,
239 confidence_threshold,
240 }
241 }
242
243 pub async fn extract_relations(&self, text: &str) -> Result<Vec<ExtractedRelation>> {
245 let sentences = self.segment_sentences(text);
247
248 let mut all_relations = Vec::new();
249
250 for sentence in sentences {
251 let entities = if self.config.enable_ner {
253 self.ner_model.extract_entities(&sentence)?
254 } else {
255 Vec::new()
256 };
257
258 let linked_entities = if self.config.enable_entity_linking {
260 self.link_entities(&entities, &sentence).await?
261 } else {
262 entities
263 };
264
265 if self.config.enable_relation_classification {
267 let relations =
268 self.extract_relations_from_entities(&sentence, &linked_entities)?;
269 all_relations.extend(relations);
270 }
271 }
272
273 let filtered_relations = all_relations
275 .into_iter()
276 .filter(|r| r.confidence >= self.confidence_threshold)
277 .collect();
278
279 Ok(filtered_relations)
280 }
281
282 pub fn to_triples(&self, relations: &[ExtractedRelation]) -> Result<Vec<Triple>> {
284 let mut triples = Vec::new();
285
286 for relation in relations {
287 let subject = if let Some(kb_id) = &relation.subject.kb_id {
289 NamedNode::new(kb_id)?
290 } else {
291 NamedNode::new(format!(
293 "http://example.org/entity/{}",
294 relation.subject.text.replace(' ', "_")
295 ))?
296 };
297
298 let predicate = NamedNode::new(format!(
300 "http://example.org/relation/{}",
301 relation.predicate.replace(' ', "_")
302 ))?;
303
304 let object = if let Some(kb_id) = &relation.object.kb_id {
306 crate::model::Object::NamedNode(NamedNode::new(kb_id)?)
307 } else {
308 match relation.object.entity_type {
310 EntityType::Date
311 | EntityType::Time
312 | EntityType::Money
313 | EntityType::Percent => {
314 crate::model::Object::Literal(Literal::new(&relation.object.text))
315 }
316 _ => crate::model::Object::NamedNode(NamedNode::new(format!(
317 "http://example.org/entity/{}",
318 relation.object.text.replace(' ', "_")
319 ))?),
320 }
321 };
322
323 let triple = Triple::new(subject, predicate, object);
324 triples.push(triple);
325 }
326
327 Ok(triples)
328 }
329
330 fn segment_sentences(&self, text: &str) -> Vec<String> {
332 text.split(". ")
334 .map(|s| s.trim().to_string())
335 .filter(|s| !s.is_empty())
336 .collect()
337 }
338
339 async fn link_entities(
341 &self,
342 entities: &[ExtractedEntity],
343 context: &str,
344 ) -> Result<Vec<ExtractedEntity>> {
345 let mut linked_entities = Vec::new();
346
347 for entity in entities {
348 let mut linked_entity = entity.clone();
349
350 if let Some(kb_id) = self.entity_linker.link_entity(entity, context)? {
353 linked_entity.kb_id = Some(kb_id);
354 }
355
356 linked_entities.push(linked_entity);
357 }
358
359 Ok(linked_entities)
360 }
361
362 fn extract_relations_from_entities(
364 &self,
365 sentence: &str,
366 entities: &[ExtractedEntity],
367 ) -> Result<Vec<ExtractedRelation>> {
368 let mut relations = Vec::new();
369
370 for (i, subject) in entities.iter().enumerate() {
372 for (j, object) in entities.iter().enumerate() {
373 if i != j {
374 if let Some((relation_type, confidence)) = self
377 .relation_model
378 .classify_relation(sentence, subject, object)?
379 {
380 let mut metadata = HashMap::new();
381 metadata.insert(
382 "extraction_method".to_string(),
383 "heuristic-keyword-match".to_string(),
384 );
385 let relation = ExtractedRelation {
386 subject: subject.clone(),
387 predicate: relation_type,
388 object: object.clone(),
389 confidence,
390 source_span: TextSpan {
391 start: 0,
392 end: sentence.len(),
393 text: sentence.to_string(),
394 },
395 context: sentence.to_string(),
396 metadata,
397 };
398
399 relations.push(relation);
400 }
401 }
402 }
403 }
404
405 Ok(relations)
406 }
407}
408
409const ORG_GAZETTEER: &[&str] = &[
411 "Inc",
412 "Inc.",
413 "Corp",
414 "Corp.",
415 "Corporation",
416 "Ltd",
417 "Ltd.",
418 "LLC",
419 "Company",
420 "GmbH",
421 "Microsoft",
422 "Google",
423 "Apple",
424 "Amazon",
425 "IBM",
426 "Oracle",
427 "Meta",
428 "Intel",
429 "Nvidia",
430];
431
432const LOCATION_GAZETTEER: &[&str] = &[
434 "Seattle",
435 "London",
436 "Paris",
437 "Tokyo",
438 "Berlin",
439 "Washington",
440 "California",
441 "France",
442 "Germany",
443 "Japan",
444 "China",
445 "India",
446 "Boston",
447 "Chicago",
448 "Amsterdam",
449 "Madrid",
450];
451
452struct HeuristicNer;
461
462impl HeuristicNer {
463 fn new() -> Self {
464 Self
465 }
466
467 fn classify_token(token: &str) -> (EntityType, f32) {
468 if LOCATION_GAZETTEER
469 .iter()
470 .any(|g| g.eq_ignore_ascii_case(token))
471 {
472 (EntityType::Location, 0.7)
473 } else if ORG_GAZETTEER.iter().any(|g| g.eq_ignore_ascii_case(token)) {
474 (EntityType::Organization, 0.7)
475 } else {
476 (EntityType::Other("Unknown".to_string()), 0.5)
478 }
479 }
480}
481
482impl NamedEntityRecognizer for HeuristicNer {
483 fn extract_entities(&self, text: &str) -> Result<Vec<ExtractedEntity>> {
484 let mut entities = Vec::new();
485 let mut search_start = 0usize;
486
487 for word in text.split_whitespace() {
488 let start = match text[search_start..].find(word) {
490 Some(rel) => search_start + rel,
491 None => continue,
492 };
493 search_start = start + word.len();
494
495 let trimmed = word.trim_matches(|c: char| !c.is_alphanumeric());
497 if trimmed.is_empty() {
498 continue;
499 }
500 let first = trimmed.chars().next().unwrap_or(' ');
501 if !first.is_uppercase() {
502 continue;
503 }
504
505 let inner_offset = word.find(trimmed).unwrap_or(0);
507 let token_start = start + inner_offset;
508 let token_end = token_start + trimmed.len();
509
510 let (entity_type, confidence) = Self::classify_token(trimmed);
511 entities.push(ExtractedEntity {
512 text: trimmed.to_string(),
513 entity_type,
514 kb_id: None,
515 confidence,
516 span: TextSpan {
517 start: token_start,
518 end: token_end,
519 text: trimmed.to_string(),
520 },
521 });
522 }
523
524 Ok(entities)
525 }
526
527 fn supported_types(&self) -> Vec<EntityType> {
528 vec![
529 EntityType::Organization,
530 EntityType::Location,
531 EntityType::Other("Unknown".to_string()),
532 ]
533 }
534}
535
536struct HeuristicRelationClassifier;
543
544impl HeuristicRelationClassifier {
545 fn new() -> Self {
546 Self
547 }
548}
549
550impl RelationClassifier for HeuristicRelationClassifier {
551 fn classify_relation(
552 &self,
553 text: &str,
554 _subject: &ExtractedEntity,
555 _object: &ExtractedEntity,
556 ) -> Result<Option<(String, f32)>> {
557 if text.contains("work") || text.contains("employ") {
558 Ok(Some(("worksFor".to_string(), 0.75)))
559 } else if text.contains("live") || text.contains("reside") {
560 Ok(Some(("livesIn".to_string(), 0.75)))
561 } else if text.contains("born") || text.contains("birth") {
562 Ok(Some(("bornIn".to_string(), 0.75)))
563 } else {
564 Ok(None)
565 }
566 }
567
568 fn supported_relations(&self) -> Vec<String> {
569 vec![
570 "worksFor".to_string(),
571 "livesIn".to_string(),
572 "bornIn".to_string(),
573 ]
574 }
575}
576
577struct LocalEntityLinker;
586
587impl LocalEntityLinker {
588 fn new() -> Self {
589 Self
590 }
591}
592
593impl EntityLinker for LocalEntityLinker {
594 fn link_entity(&self, _entity: &ExtractedEntity, _context: &str) -> Result<Option<String>> {
595 Ok(None)
597 }
598
599 fn kb_info(&self) -> KnowledgeBaseInfo {
600 KnowledgeBaseInfo {
601 name: "none (no knowledge base configured)".to_string(),
602 base_uri: String::new(),
603 version: "n/a".to_string(),
604 entity_count: 0,
605 }
606 }
607}
608
609#[cfg(test)]
610mod tests {
611 use super::*;
612 use crate::ai::AiConfig;
613
614 #[tokio::test]
615 async fn test_relation_extractor_creation() {
616 let config = AiConfig::default();
617 let extractor = RelationExtractor::new(&config);
618 assert!(extractor.is_ok());
619 }
620
621 #[tokio::test]
622 async fn test_relation_extraction() {
623 let config = AiConfig::default();
624 let extractor = RelationExtractor::new(&config).expect("construction should succeed");
625
626 let text = "John works for Microsoft. He lives in Seattle.";
627 let relations = extractor
628 .extract_relations(text)
629 .await
630 .expect("async operation should succeed");
631
632 assert!(!relations.is_empty());
634 }
635
636 #[test]
637 fn test_sentence_segmentation() {
638 let config = AiConfig::default();
639 let extractor = RelationExtractor::new(&config).expect("construction should succeed");
640
641 let text = "First sentence. Second sentence. Third sentence.";
642 let sentences = extractor.segment_sentences(text);
643
644 assert_eq!(sentences.len(), 3);
645 assert_eq!(sentences[0], "First sentence");
646 }
647
648 #[test]
649 fn test_to_triples() {
650 let config = AiConfig::default();
651 let extractor = RelationExtractor::new(&config).expect("construction should succeed");
652
653 let relation = ExtractedRelation {
654 subject: ExtractedEntity {
655 text: "John".to_string(),
656 entity_type: EntityType::Person,
657 kb_id: None,
658 confidence: 0.9,
659 span: TextSpan {
660 start: 0,
661 end: 4,
662 text: "John".to_string(),
663 },
664 },
665 predicate: "worksFor".to_string(),
666 object: ExtractedEntity {
667 text: "Microsoft".to_string(),
668 entity_type: EntityType::Organization,
669 kb_id: None,
670 confidence: 0.85,
671 span: TextSpan {
672 start: 15,
673 end: 24,
674 text: "Microsoft".to_string(),
675 },
676 },
677 confidence: 0.8,
678 source_span: TextSpan {
679 start: 0,
680 end: 25,
681 text: "John works for Microsoft.".to_string(),
682 },
683 context: "John works for Microsoft.".to_string(),
684 metadata: HashMap::new(),
685 };
686
687 let triples = extractor
688 .to_triples(&[relation])
689 .expect("operation should succeed");
690 assert_eq!(triples.len(), 1);
691 }
692
693 #[test]
694 fn regression_entity_linker_does_not_fabricate_dbpedia_uris() {
695 let linker = LocalEntityLinker::new();
696 let entity = ExtractedEntity {
697 text: "John".to_string(),
698 entity_type: EntityType::Person,
699 kb_id: None,
700 confidence: 0.5,
701 span: TextSpan {
702 start: 0,
703 end: 4,
704 text: "John".to_string(),
705 },
706 };
707 let linked = linker.link_entity(&entity, "context").expect("link");
709 assert_eq!(linked, None);
710
711 let info = linker.kb_info();
713 assert_eq!(info.entity_count, 0);
714 assert!(!info.base_uri.contains("dbpedia"));
715 }
716
717 #[test]
718 fn regression_ner_reports_real_byte_offsets() {
719 let ner = HeuristicNer::new();
720 let text = "John works for Microsoft";
721 let entities = ner.extract_entities(text).expect("ner");
722
723 assert!(!entities.is_empty());
725 for entity in &entities {
726 assert_eq!(&text[entity.span.start..entity.span.end], entity.span.text);
727 assert_eq!(entity.span.text, entity.text);
728 }
729
730 let microsoft = entities
733 .iter()
734 .find(|e| e.text == "Microsoft")
735 .expect("Microsoft detected");
736 assert_eq!(microsoft.span.start, 15);
737 assert!(matches!(microsoft.entity_type, EntityType::Organization));
738
739 let john = entities
741 .iter()
742 .find(|e| e.text == "John")
743 .expect("John detected");
744 assert!(matches!(john.entity_type, EntityType::Other(_)));
745 }
746
747 #[tokio::test]
748 async fn regression_extracted_relations_tagged_as_heuristic() {
749 let config = AiConfig::default();
750 let extractor = RelationExtractor::new(&config).expect("construction");
751 let relations = extractor
752 .extract_relations("John works for Microsoft")
753 .await
754 .expect("extract");
755 assert!(!relations.is_empty());
756 for relation in &relations {
757 assert_eq!(
758 relation
759 .metadata
760 .get("extraction_method")
761 .map(String::as_str),
762 Some("heuristic-keyword-match")
763 );
764 }
765 }
766}