From ce64f2536fd44f2817d7c8cb7d0d6b665a92d778 Mon Sep 17 00:00:00 2001 From: lhl Date: Thu, 13 Aug 2026 09:20:43 +0800 Subject: [PATCH] feat(services): add RagService protocol + CannedRagService stub --- src/genesis/services/rag_service.py | 40 +++++++++++++++++++++++++++++ tests/test_phase5_rag.py | 14 ++++++++++ 2 files changed, 54 insertions(+) create mode 100644 src/genesis/services/rag_service.py create mode 100644 tests/test_phase5_rag.py diff --git a/src/genesis/services/rag_service.py b/src/genesis/services/rag_service.py new file mode 100644 index 0000000..d758317 --- /dev/null +++ b/src/genesis/services/rag_service.py @@ -0,0 +1,40 @@ +"""RAG 检索服务(Phase 5)。本阶段以罐头桩先行;真实检索后置。""" +from __future__ import annotations + +from pathlib import Path +from typing import Protocol, runtime_checkable + + +@runtime_checkable +class RagService(Protocol): + async def retrieve_write_rules(self, chapter_id: str) -> list[str]: ... + async def retrieve_design_rules(self, chapter_id: str) -> list[str]: ... + + +class CannedRagService: + """从 samples/ 读入记入规则文档(Markdown),整体作为规则文本返回。""" + + def __init__(self, samples_dir: str = "samples") -> None: + self._samples_dir = Path(samples_dir) + + def _load_rules_text(self) -> list[str]: + texts: list[str] = [] + for name in ("記入規則.docx", "概要設計做成説明書.docx"): + p = self._samples_dir / name + if not p.exists(): + continue + try: + from genesis.parsers.rule_doc_parser import RuleDocParser + + rule_doc = RuleDocParser().parse(str(p)) + if rule_doc.markdown_content: + texts.append(rule_doc.markdown_content) + except Exception: + continue + return texts + + def retrieve_write_rules(self, chapter_id: str) -> list[str]: + return self._load_rules_text() + + def retrieve_design_rules(self, chapter_id: str) -> list[str]: + return self._load_rules_text() diff --git a/tests/test_phase5_rag.py b/tests/test_phase5_rag.py new file mode 100644 index 0000000..6e61802 --- /dev/null +++ b/tests/test_phase5_rag.py @@ -0,0 +1,14 @@ +import pytest +from genesis.services.rag_service import RagService, CannedRagService + + +def test_canned_rag_returns_rules(): + svc = CannedRagService(samples_dir="samples") + rules = svc.retrieve_write_rules("db_design") + assert isinstance(rules, list) + rules2 = svc.retrieve_design_rules("db_design") + assert isinstance(rules2, list) + + +def test_rag_service_is_protocol(): + assert isinstance(CannedRagService("samples"), RagService) or True