From eee9932a4e4bda9eb5bcb60d201fcec121a04b48 Mon Sep 17 00:00:00 2001 From: Axy Date: Thu, 27 Aug 2026 21:10:54 +0200 Subject: [PATCH] Refactoring a little --- src/rag/__init__.py | 127 +++++++++++++++++++++----------------------- 1 file changed, 61 insertions(+), 66 deletions(-) diff --git a/src/rag/__init__.py b/src/rag/__init__.py index 8dab130..f0b6e63 100644 --- a/src/rag/__init__.py +++ b/src/rag/__init__.py @@ -1,5 +1,6 @@ import logging import os +from collections.abc import Callable from sys import stderr import fire @@ -34,6 +35,29 @@ class RAG: self._bm25_store: BM25[Source] | None = None self._inference_store: InferenceBackend | None = None + @staticmethod + def _readf[T](path: str, cb: Callable[[str], T], name: str = "file") -> T: + try: + with open(path) as f: + return cb(f.read()) + except Exception as e: + logging.getLogger(__name__).error(f"Failed to read {name}: {e}") + exit(1) + + @staticmethod + def _writef(path: str, data: bytes | str, name: str = "file") -> None: + try: + os.makedirs(os.path.dirname(path), exist_ok=True) + if isinstance(data, bytes): + with open(path, "wb") as f: + f.write(data) + else: + with open(path, "w") as f: + f.write(data) + except Exception as e: + logging.getLogger(__name__).error(f"Failed to write {name}: {e}") + exit(1) + def index( self, /, @@ -57,42 +81,31 @@ class RAG: bm25 = BM25.from_corpus( tqdm(chunks.items(), desc="Indexing files"), stopwords("en") ) - try: - os.makedirs(index, exist_ok=True) - with open(index + "/index.json", "wb") as f: - print(f"Writing index to {index}...", file=stderr) - f.write( - storage_adapter.dump_json( - { - key_adapter.dump_json(k): v - for k, v in bm25.to_storage() - } - ) - ) - except OSError as e: - logging.getLogger(__name__).error( - f"Failed to write index file {e}" - ) - exit(1) - print(f"Successfully ingested: {data} -> {index}") + self._writef( + index + "/index.json", + storage_adapter.dump_json( + {key_adapter.dump_json(k): v for k, v in bm25.to_storage()} + ), + name="index", + ) + print(f"Successfully ingested: {data} into {index}") def _bm25(self, index: str) -> BM25[Source]: if self._bm25_store: return self._bm25_store - try: - with open(index + "/index.json") as f: - bm25: BM25[Source] = BM25.from_storage( - (key_adapter.validate_json(k), v) - for k, v in tqdm( - storage_adapter.validate_json(f.read()).items(), - desc="Unpacking index", - ) + bm25: BM25[Source] = self._readf( + index + "/index.json", + lambda s: BM25.from_storage( + (key_adapter.validate_json(k), v) + for k, v in tqdm( + storage_adapter.validate_json(s).items(), + desc="Unpacking index", ) - self._bm25_store = bm25 - return bm25 - except (OSError, ValidationError) as e: - logging.getLogger(__name__).error(f"Failed to open index {e}") - exit(1) + ), + name="index", + ) + self._bm25_store = bm25 + return bm25 def _search(self, query: str, k: int, index: str) -> list[Source]: bm25 = self._bm25(index) @@ -113,12 +126,7 @@ class RAG: index: str = "data/processed", ) -> None: bm25 = self._bm25(index) - try: - with open(dataset_path) as f: - dataset = RagDataset.model_validate_json(f.read()) - except (OSError, ValidationError) as e: - logging.getLogger(__name__).error(f"Failed to open dataset {e}") - exit(1) + dataset = self._readf(dataset_path, RagDataset.model_validate_json) result = SearchResults( search_results=[ RetrievedQuestion( @@ -132,15 +140,11 @@ class RAG: ], k=k, ) - try: - os.makedirs(save_directory, exist_ok=True) - with open( - save_directory + "/" + os.path.basename(dataset_path), "w" - ) as f: - f.write(result.model_dump_json()) - except (OSError, ValidationError) as e: - logging.getLogger(__name__).error(f"Failed to open dataset {e}") - exit(1) + self._writef( + save_directory + "/" + os.path.basename(dataset_path), + result.model_dump_json(), + name="dataset", + ) def _inference(self) -> InferenceBackend: if not self._inference_store: @@ -176,14 +180,11 @@ class RAG: max_tokens: int = 100, ) -> None: bm25 = self._bm25(index) - try: - with open(student_search_result_path) as f: - dataset = SearchResults.model_validate_json(f.read()) - except (OSError, ValidationError) as e: - logging.getLogger(__name__).error( - f"Failed to open search results {e}" - ) - exit(1) + dataset = self._readf( + student_search_result_path, + SearchResults.model_validate_json, + name="search results", + ) inference = self._inference() result = SearchAnswers( search_results=[ @@ -194,15 +195,9 @@ class RAG: ], k=dataset.k, ) - try: - os.makedirs(save_directory, exist_ok=True) - with open( - save_directory - + "/" - + os.path.basename(student_search_result_path), - "w", - ) as f: - f.write(result.model_dump_json()) - except (OSError, ValidationError) as e: - logging.getLogger(__name__).error(f"Failed to open dataset {e}") - exit(1) + self._writef( + save_directory + + "/" + + os.path.basename(student_search_result_path), + result.model_dump_json(), + ) -- 2.53.0