]> Untitled Git - axy/ft/rag.git/commitdiff
Refactoring a little
authorAxy <gilliardmarthey.axel@gmail.com>
Thu, 27 Aug 2026 19:10:54 +0000 (21:10 +0200)
committerAxy <gilliardmarthey.axel@gmail.com>
Thu, 27 Aug 2026 19:10:54 +0000 (21:10 +0200)
src/rag/__init__.py

index 8dab130adcc7407a642eb0948edfdcdc5aad1690..f0b6e631b97c9c7ef63e13348176793c02e643c6 100644 (file)
@@ -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(),
+        )