1use std::collections::HashMap;
11
12use color_eyre::Result;
13use polars::prelude::*;
14
15use crate::formats::model_files::MetaValue;
16use crate::formats::text_formats::{Detail, capped_list, count, note};
17use crate::notes::Note;
18
19pub(crate) const READER: crate::formats::readers::Reader = crate::formats::readers::Reader {
21 convert: Some(|input| {
22 crate::formats::text_formats::convert_with(input, SdfReader::new(), |reader, lf| {
23 if reader.stats().records == 0 {
24 return Err(color_eyre::eyre::eyre!("no SDF records"));
25 }
26 let lf = type_fields(lf, reader.fields());
27 Ok((lf, notes(reader.stats()), detail(reader)))
28 })
29 }),
30 scan: crate::formats::readers::read_into,
31 signatures: &[crate::formats::readers::Signature {
32 says: |head, _| looks_like(head),
33 kind: crate::formats::readers::Kind::Text,
34 trusted: crate::formats::readers::Trusted {
35 listing: false,
36 ..crate::formats::readers::EVERYWHERE
37 },
38 }],
39 ..crate::formats::readers::BASE
40};
41
42pub const MAX_LINE: usize = 1 << 20;
44pub const MAX_VALUE: usize = 1 << 20;
46pub const MAX_NAME: usize = 256;
48pub const BATCH_ROWS: usize = 16_384;
50pub const BATCH_TEXT: usize = 32 << 20;
54
55pub const CORE: [&str; 3] = ["name", "atoms", "bonds"];
57
58pub fn looks_like(head: &[u8]) -> bool {
61 let text = head.strip_prefix(b"\xef\xbb\xbf").unwrap_or(head);
62 let mut lines = text.split(|&b| b == b'\n');
63 let counts = lines.nth(3).unwrap_or_default().trim_ascii();
65 let counts_line = (counts.ends_with(b"V2000") || counts.ends_with(b"V3000"))
66 && counts[..counts.len() - 5]
67 .iter()
68 .all(|b| b.is_ascii_digit() || b.is_ascii_whitespace());
69 let has = |needle: &[u8]| text.windows(needle.len()).any(|w| w == needle);
70 counts_line || (has(b"M END") && (has(b"$$$$") || has(b"> <")))
71}
72
73#[derive(Debug, Clone, PartialEq)]
75pub struct FieldColumn {
76 pub name: String,
77 pub values: u64,
79 pub integers: bool,
81 pub numbers: bool,
83}
84
85impl FieldColumn {
86 pub fn dtype(&self) -> DataType {
88 if self.values == 0 {
89 DataType::String
90 } else if self.integers {
91 DataType::Int64
92 } else if self.numbers {
93 DataType::Float64
94 } else {
95 DataType::String
96 }
97 }
98}
99
100#[derive(Debug, Clone, Default, PartialEq)]
102pub struct Stats {
103 pub records: u64,
104 pub long_lines: u64,
106 pub long_values: u64,
108 pub repeated: u64,
110 pub fields_dropped: u64,
112 pub v3000: u64,
114 pub unterminated: bool,
116}
117
118#[derive(Debug, Clone, Copy, PartialEq, Eq)]
119enum Place {
120 Header(u8),
122 Ctab,
124 Data,
126 Value,
128}
129
130#[derive(Debug, Default)]
132struct Record {
133 name: Option<String>,
134 atoms: Option<u32>,
135 bonds: Option<u32>,
136 values: HashMap<usize, String>,
137 field: Option<usize>,
139 value: String,
140 cut: bool,
142 started: bool,
144}
145
146#[derive(Debug)]
148pub struct SdfReader {
149 buf: Vec<u8>,
150 cutting: bool,
152 place: Place,
153 record: Record,
154 fields: Vec<FieldColumn>,
155 by_name: HashMap<String, usize>,
156 names: Vec<Option<String>>,
157 atoms: Vec<Option<u32>>,
158 bonds: Vec<Option<u32>>,
159 columns: Vec<Vec<Option<String>>>,
161 rows: usize,
162 held: usize,
163 stats: Stats,
164 max_fields: usize,
166}
167
168impl Default for SdfReader {
169 fn default() -> Self {
170 Self::new()
171 }
172}
173
174impl SdfReader {
175 pub fn new() -> Self {
176 Self {
177 buf: Vec::new(),
178 cutting: false,
179 place: Place::Header(0),
180 record: Record::default(),
181 fields: Vec::new(),
182 by_name: HashMap::new(),
183 names: Vec::new(),
184 atoms: Vec::new(),
185 bonds: Vec::new(),
186 columns: Vec::new(),
187 rows: 0,
188 held: 0,
189 stats: Stats::default(),
190 max_fields: crate::limits::get().sdf_fields,
191 }
192 }
193
194 pub fn stats(&self) -> &Stats {
195 &self.stats
196 }
197
198 pub fn fields(&self) -> &[FieldColumn] {
200 &self.fields
201 }
202
203 pub fn push(&mut self, bytes: &[u8]) {
205 let mut rest = bytes;
206 while let Some(end) = rest.iter().position(|&b| b == b'\n') {
207 self.take(&rest[..end]);
208 if !self.cutting {
209 let line = std::mem::take(&mut self.buf);
210 self.line(&line);
211 }
212 self.cutting = false;
213 self.buf.clear();
214 rest = &rest[end + 1..];
215 }
216 self.take(rest);
217 }
218
219 fn take(&mut self, part: &[u8]) {
221 if self.cutting {
222 return;
223 }
224 let room = MAX_LINE.saturating_sub(self.buf.len());
225 if part.len() > room {
226 self.buf.extend_from_slice(&part[..room]);
227 self.stats.long_lines += 1;
228 let line = std::mem::take(&mut self.buf);
229 self.line(&line);
230 self.cutting = true;
231 } else {
232 self.buf.extend_from_slice(part);
233 }
234 }
235
236 pub fn take_batch(&mut self) -> PolarsResult<Option<DataFrame>> {
238 if self.rows < BATCH_ROWS && self.held < BATCH_TEXT {
239 return Ok(None);
240 }
241 self.batch().map(Some)
242 }
243
244 pub fn finish(&mut self) -> PolarsResult<DataFrame> {
246 if !self.cutting && !self.buf.is_empty() {
247 let line = std::mem::take(&mut self.buf);
248 self.line(&line);
249 }
250 self.buf.clear();
251 self.cutting = false;
252 if self.record.started {
253 self.stats.unterminated = true;
254 self.end_record();
255 }
256 self.batch()
257 }
258
259 fn batch(&mut self) -> PolarsResult<DataFrame> {
260 self.held = 0;
261 let height = std::mem::take(&mut self.rows);
262 let mut columns = vec![
263 StringChunked::from_iter_options(
264 "name".into(),
265 std::mem::take(&mut self.names).into_iter(),
266 )
267 .into_column(),
268 UInt32Chunked::from_iter_options(
269 "atoms".into(),
270 std::mem::take(&mut self.atoms).into_iter(),
271 )
272 .into_column(),
273 UInt32Chunked::from_iter_options(
274 "bonds".into(),
275 std::mem::take(&mut self.bonds).into_iter(),
276 )
277 .into_column(),
278 ];
279 for (field, values) in self.fields.iter().zip(self.columns.iter_mut()) {
280 columns.push(
281 StringChunked::from_iter_options(
282 field.name.as_str().into(),
283 std::mem::take(values).into_iter(),
284 )
285 .into_column(),
286 );
287 }
288 DataFrame::new(height, columns)
289 }
290
291 fn line(&mut self, raw: &[u8]) {
292 let raw = raw.strip_suffix(b"\r").unwrap_or(raw);
293 let line = String::from_utf8_lossy(raw);
294 let line = line.as_ref();
295 if line.starts_with("$$$$") {
296 self.end_value();
297 if self.record.started {
298 self.end_record();
299 } else {
300 self.record = Record::default();
302 self.place = Place::Header(0);
303 }
304 return;
305 }
306 if !line.trim().is_empty() {
307 self.record.started = true;
308 }
309 match self.place {
310 Place::Header(n) => {
311 if n == 0 {
313 self.record.name = Some(line.trim().to_string()).filter(|s| !s.is_empty());
314 } else if n == 3 {
315 self.counts(line);
316 }
317 self.place = if n >= 3 {
318 Place::Ctab
319 } else {
320 Place::Header(n + 1)
321 };
322 }
323 Place::Ctab => {
324 if line.starts_with("M END") {
325 self.place = Place::Data;
326 } else if let Some(counts) = line.strip_prefix("M V30 COUNTS") {
327 let mut numbers = counts.split_whitespace();
328 self.record.atoms = numbers.next().and_then(|n| n.parse().ok());
329 self.record.bonds = numbers.next().and_then(|n| n.parse().ok());
330 } else if is_item_header(line) {
331 self.place = Place::Data;
333 self.item(line);
334 }
335 }
336 Place::Data => {
337 if is_item_header(line) {
338 self.item(line);
339 }
340 }
341 Place::Value => {
342 if line.trim().is_empty() {
343 self.end_value();
344 self.place = Place::Data;
345 } else if is_item_header(line) && line.contains('<') {
346 self.end_value();
348 self.item(line);
349 } else if self.record.field.is_some() && !self.record.cut {
350 let value = &mut self.record.value;
351 let sep = usize::from(!value.is_empty());
352 let room = MAX_VALUE.saturating_sub(value.len() + sep);
353 let mut take = line.len().min(room);
354 while !line.is_char_boundary(take) {
355 take -= 1;
356 }
357 if take > 0 {
358 if sep == 1 {
359 value.push('\n');
360 }
361 value.push_str(&line[..take]);
362 }
363 if take < line.len() {
364 self.stats.long_values += 1;
365 self.record.cut = true;
366 }
367 }
368 }
369 }
370 }
371
372 fn counts(&mut self, line: &str) {
374 if line.contains("V3000") {
375 self.stats.v3000 += 1;
376 return;
377 }
378 let number = |range: std::ops::Range<usize>| {
379 line.get(range).and_then(|s| s.trim().parse::<u32>().ok())
380 };
381 self.record.atoms = number(0..3);
382 self.record.bonds = number(3..6);
383 }
384
385 fn item(&mut self, line: &str) {
387 self.place = Place::Value;
388 self.record.value.clear();
389 self.record.field = None;
390 self.record.cut = false;
391 let Some(name) = item_name(line) else {
392 return;
393 };
394 let index = match self.by_name.get(&name) {
395 Some(&i) => i,
396 None if self.fields.len() >= self.max_fields => {
397 self.stats.fields_dropped += 1;
398 return;
399 }
400 None => {
401 let i = self.fields.len();
402 self.by_name.insert(name.clone(), i);
403 let mut column = name.clone();
406 let mut n = 1;
407 while CORE.contains(&column.as_str())
408 || self.fields.iter().any(|f| f.name == column)
409 {
410 n += 1;
411 column = format!("{name}_{n}");
412 }
413 self.fields.push(FieldColumn {
414 name: column,
415 values: 0,
416 integers: true,
417 numbers: true,
418 });
419 self.columns.push(vec![None; self.rows]);
420 i
421 }
422 };
423 if self.record.values.contains_key(&index) {
424 self.stats.repeated += 1;
425 return;
426 }
427 self.record.field = Some(index);
428 }
429
430 fn end_value(&mut self) {
431 if self.place != Place::Value {
432 return;
433 }
434 let value = std::mem::take(&mut self.record.value);
435 if let Some(field) = self.record.field.take() {
436 self.keep(field, value);
437 }
438 }
439
440 fn keep(&mut self, field: usize, value: String) {
441 let trimmed = value.trim();
442 if trimmed.is_empty() {
443 return;
444 }
445 let column = &mut self.fields[field];
446 column.values += 1;
447 column.integers &= trimmed.parse::<i64>().is_ok();
448 column.numbers &= trimmed.parse::<f64>().is_ok_and(f64::is_finite);
449 let value = if trimmed.len() == value.len() {
450 value
451 } else {
452 trimmed.to_string()
453 };
454 self.record.values.insert(field, value);
455 }
456
457 fn end_record(&mut self) {
458 let record = std::mem::take(&mut self.record);
459 self.place = Place::Header(0);
460 self.names.push(record.name);
461 self.atoms.push(record.atoms);
462 self.bonds.push(record.bonds);
463 let mut values = record.values;
464 for (i, column) in self.columns.iter_mut().enumerate() {
465 let value = values.remove(&i);
466 self.held += size_of::<Option<String>>() + value.as_ref().map_or(0, String::len);
467 column.push(value);
468 }
469 self.rows += 1;
470 self.stats.records += 1;
471 }
472}
473
474fn is_item_header(line: &str) -> bool {
476 line.starts_with('>')
477}
478
479pub fn item_name(line: &str) -> Option<String> {
482 let rest = line.strip_prefix('>')?;
483 let name = match rest.find('<') {
484 Some(open) => {
485 let after = &rest[open + 1..];
486 let close = after.find('>')?;
487 &after[..close]
488 }
489 None => rest.split_whitespace().next()?,
490 };
491 let name = name.trim();
492 if name.is_empty() {
493 return None;
494 }
495 let mut cut = name.len().min(MAX_NAME);
496 while !name.is_char_boundary(cut) {
497 cut -= 1;
498 }
499 Some(name[..cut].to_string())
500}
501
502fn type_fields(lf: LazyFrame, fields: &[FieldColumn]) -> LazyFrame {
504 let casts: Vec<Expr> = fields
505 .iter()
506 .filter(|f| f.dtype() != DataType::String)
507 .map(|f| col(f.name.as_str()).cast(f.dtype()))
508 .collect();
509 if casts.is_empty() {
510 lf
511 } else {
512 lf.with_columns(casts)
513 }
514}
515
516fn notes(stats: &Stats) -> Vec<Note> {
517 let mut notes = Vec::new();
518 let of_records = format!("of {}", count(stats.records, "record", "records"));
519 if stats.long_lines > 0 {
520 notes.push(note(
521 format!(
522 "{} cut at {} MiB",
523 count(stats.long_lines, "line", "lines"),
524 MAX_LINE >> 20
525 ),
526 "in the whole file".to_string(),
527 ));
528 }
529 if stats.long_values > 0 {
530 notes.push(note(
531 format!(
532 "{} cut at {} MiB",
533 count(stats.long_values, "value", "values"),
534 MAX_VALUE >> 20
535 ),
536 of_records.clone(),
537 ));
538 }
539 if stats.repeated > 0 {
540 notes.push(note(
541 format!(
542 "repeated field: {} {} first kept",
543 count(stats.repeated, "data item", "data items"),
544 crate::glyphs::get().middot
545 ),
546 of_records.clone(),
547 ));
548 }
549 if stats.fields_dropped > 0 {
550 notes.push(note(
551 crate::limits::left_out(
552 &count(stats.fields_dropped, "data item", "data items"),
553 crate::limits::get().sdf_fields,
554 "sdf_fields",
555 ),
556 of_records.clone(),
557 ));
558 }
559 if stats.unterminated {
560 notes.push(note(
561 format!(
562 "last record has no $$$$ {} read as is",
563 crate::glyphs::get().middot
564 ),
565 of_records,
566 ));
567 }
568 notes
569}
570
571pub fn detail(reader: &SdfReader) -> Detail {
573 let stats = reader.stats();
574 let sep = format!(" {} ", crate::glyphs::get().middot);
575 let mut head = format!(
576 "SDF{sep}{}{sep}{}",
577 count(stats.records, "record", "records"),
578 count(reader.fields().len() as u64, "field", "fields")
579 );
580 if stats.v3000 > 0 {
581 head.push_str(&sep);
582 head.push_str(&format!(
583 "{} V3000",
584 count(stats.v3000, "record", "records")
585 ));
586 }
587 let list = capped_list(
588 reader.fields().iter().map(|f| {
589 (
590 f.name.clone(),
591 MetaValue::Text(format!(
592 "{}{sep}in {} of {}",
593 f.dtype(),
594 crate::numfmt::group_chrome(f.values as usize),
595 count(stats.records, "record", "records"),
596 )),
597 )
598 }),
599 reader.fields().len(),
600 );
601 Detail {
602 tab: crate::formats::text_formats::tab(crate::FileFormat::Sdf),
603 lines: vec![head],
604 list_title: "Fields",
605 list,
606 first: false,
607 ..Default::default()
608 }
609}
610
611impl crate::formats::text_formats::BatchReader for SdfReader {
612 fn push(&mut self, piece: &[u8]) -> Result<()> {
613 self.push(piece);
614 Ok(())
615 }
616
617 fn take_batch(&mut self) -> PolarsResult<Option<DataFrame>> {
618 self.take_batch()
619 }
620
621 fn finish(&mut self) -> Result<DataFrame> {
622 Ok(self.finish()?)
623 }
624}
625
626#[cfg(test)]
627mod tests {
628 use super::*;
629
630 #[test]
632 fn errors_name_the_file() {
633 crate::formats::readers::bad_input::each_names_its_file(
634 crate::FileFormat::Sdf,
635 &[("empty.sdf", b"", "No SDF records")],
636 );
637 }
638
639 const SAMPLE: &str = "aspirin
640 RDKit 2D
641
642 3 2 0 0 0 0 0 0 0 0999 V2000
643 0.0000 0.0000 0.0000 C 0 0 0 0 0 0 0 0 0 0 0 0
644 1.2990 0.7500 0.0000 C 0 0
645 2.5981 -0.0000 0.0000 O 0 0
646 1 2 1 0
647 2 3 2 0
648M END
649> <ID>
650101
651
652> <LogP> (MD-1)
6531.19
654
655> <SMILES>
656CC(=O)O
657
658$$$$
659
660 RDKit 2D
661
662 1 0 0 0 0 0 0 0 0 0999 V2000
663 0.0000 0.0000 0.0000 N 0 0
664M END
665> <ID>
666102
667
668> 25 <Notes>
669first line
670second line
671
672> <LogP>
673n/a
674
675$$$$
676";
677
678 fn read(text: &[u8], piece: usize) -> (DataFrame, SdfReader) {
679 let mut reader = SdfReader::new();
680 let mut frames = Vec::new();
681 for chunk in text.chunks(piece) {
682 reader.push(chunk);
683 if let Some(df) = reader.take_batch().unwrap() {
684 frames.push(df);
685 }
686 }
687 frames.push(reader.finish().unwrap());
688 let width = frames.last().unwrap().width();
689 let mut df = frames.pop().unwrap();
690 for f in frames.into_iter().rev() {
691 assert!(f.width() <= width);
692 df = f.vstack(&df).unwrap_or(df);
693 }
694 (df, reader)
695 }
696
697 fn strings(df: &DataFrame, name: &str) -> Vec<Option<String>> {
698 df.column(name)
699 .unwrap()
700 .str()
701 .unwrap()
702 .iter()
703 .map(|s| s.map(String::from))
704 .collect()
705 }
706
707 #[test]
710 fn a_wide_file_hands_over_small_batches() {
711 let mut reader = SdfReader::new();
712 let mut wide = String::from("wide\n\n\n 0 0 0 0 0 0 0 0 0 0999 V2000\nM END\n");
713 for i in 0..2_000 {
714 wide.push_str(&format!("> <F{i}>\nx\n\n"));
715 }
716 wide.push_str("$$$$\n");
717 reader.push(wide.as_bytes());
718 let narrow = "n\n\n\n 0 0 0 0 0 0 0 0 0 0999 V2000\nM END\n$$$$\n";
719 let mut rows_at_batch = None;
720 for row in 1..BATCH_ROWS {
721 reader.push(narrow.as_bytes());
722 if let Some(df) = reader.take_batch().unwrap() {
723 rows_at_batch = Some((row, df.height()));
724 break;
725 }
726 }
727 let (row, height) = rows_at_batch.expect("a batch is handed over");
728 assert!(height < BATCH_ROWS / 8, "{height} rows of 2,003 columns");
729 assert_eq!(height, row + 1);
730 }
731
732 #[test]
733 fn records_read_as_rows_with_their_fields() {
734 for piece in [1, 5, 64, 4096] {
735 let (df, reader) = read(SAMPLE.as_bytes(), piece);
736 assert_eq!(df.height(), 2, "piece {piece}");
737 assert_eq!(
738 df.get_column_names(),
739 ["name", "atoms", "bonds", "ID", "LogP", "SMILES", "Notes"]
740 );
741 assert_eq!(strings(&df, "name"), [Some("aspirin".into()), None]);
742 let atoms: Vec<_> = df.column("atoms").unwrap().u32().unwrap().iter().collect();
743 assert_eq!(atoms, [Some(3), Some(1)]);
744 let bonds: Vec<_> = df.column("bonds").unwrap().u32().unwrap().iter().collect();
745 assert_eq!(bonds, [Some(2), Some(0)]);
746 assert_eq!(
747 strings(&df, "Notes"),
748 [None, Some("first line\nsecond line".into())]
749 );
750 assert_eq!(strings(&df, "SMILES"), [Some("CC(=O)O".into()), None]);
751 let fields = reader.fields();
752 assert_eq!(fields[0].dtype(), DataType::Int64);
753 assert_eq!(fields[1].dtype(), DataType::String, "n/a is not a number");
754 assert_eq!(reader.stats().records, 2);
755 assert!(!reader.stats().unterminated);
756 }
757 }
758
759 #[test]
760 fn field_types_are_cast_on_the_frame() {
761 let (df, reader) = read(b"m\n\n\n 0 0 0 0 0 0 0 0 0 0999 V2000\nM END\n> <MW>\n180.16\n\n> <N>\n3\n\n$$$$\n", 7);
762 let lf = type_fields(df.lazy(), reader.fields());
763 let df = lf.collect().unwrap();
764 assert_eq!(df.column("MW").unwrap().dtype(), &DataType::Float64);
765 assert_eq!(df.column("N").unwrap().i64().unwrap().get(0), Some(3));
766 }
767
768 #[test]
769 fn v3000_counts_and_an_unterminated_record() {
770 let text = "x\n\n\n 0 0 0 0 0 999 V3000\nM V30 BEGIN CTAB\nM V30 COUNTS 12 11 0 0 0\nM V30 END CTAB\nM END\n> <A>\n1\n";
771 let (df, reader) = read(text.as_bytes(), 3);
772 assert_eq!(df.column("atoms").unwrap().u32().unwrap().get(0), Some(12));
773 assert_eq!(df.column("bonds").unwrap().u32().unwrap().get(0), Some(11));
774 assert!(reader.stats().unterminated);
775 assert_eq!(reader.stats().v3000, 1);
776 }
777
778 #[test]
779 fn item_names() {
780 assert_eq!(item_name("> <KEY>").as_deref(), Some("KEY"));
781 assert_eq!(item_name("> <KEY> (ID)").as_deref(), Some("KEY"));
782 assert_eq!(item_name("> 25 <a b>").as_deref(), Some("a b"));
783 assert_eq!(item_name("> DT12 (MD-08974)").as_deref(), Some("DT12"));
784 assert_eq!(item_name("> <>"), None);
785 assert_eq!(item_name(">"), None);
786 }
787
788 #[test]
789 fn long_lines_and_values_are_cut() {
790 let mut text = b"m\n\n\n 0 0\nM END\n> <Big>\n".to_vec();
791 text.extend(std::iter::repeat_n(b'a', MAX_LINE + 5));
792 text.extend_from_slice(b"\nmore\n\n$$$$\n");
793 let (df, reader) = read(&text, 1 << 16);
794 let big = df
795 .column("Big")
796 .unwrap()
797 .str()
798 .unwrap()
799 .get(0)
800 .unwrap()
801 .len();
802 assert!(big <= MAX_VALUE, "{big}");
803 assert_eq!(reader.stats().long_lines, 1);
804 assert_eq!(reader.stats().long_values, 1);
805 }
806
807 #[test]
808 fn sniffing() {
809 assert!(looks_like(SAMPLE.as_bytes()));
810 assert!(looks_like(b"\n\n\nM END\n> <A>\n1\n\n$$$$\n"));
811 assert!(!looks_like(b"a,b\n1,2\n"));
812 assert!(
813 !looks_like(b"model,code\nx,1\ny,2\nz,V2000\n"),
814 "a word, not a counts line"
815 );
816 }
817
818 #[test]
819 fn a_field_named_like_a_core_column_gets_its_own() {
820 let mut reader = SdfReader::new();
821 reader.push(
822 b"aspirin\n\n\n 0 0 0 0 0 0 0 0 0 0999 V2000\nM END\n\
823 > <name>\nASA\n\n> <name_2>\nx\n\n$$$$\n",
824 );
825 let df = reader.finish().unwrap();
826 let names: Vec<&str> = df.get_column_names().iter().map(|c| c.as_str()).collect();
827 assert_eq!(names, ["name", "atoms", "bonds", "name_2", "name_2_2"]);
828 assert_eq!(strings(&df, "name"), [Some("aspirin".into())]);
829 assert_eq!(strings(&df, "name_2"), [Some("ASA".into())]);
830 }
831}