diff --git a/src/genesis/writer/writer_agent.py b/src/genesis/writer/writer_agent.py new file mode 100644 index 0000000..2fbd7f1 --- /dev/null +++ b/src/genesis/writer/writer_agent.py @@ -0,0 +1,114 @@ +"""Writer Agent:调用推理引擎生成单章内容(Phase 5)。 + +接入真实 InferenceEngine.chat_structured(session_id/prompt/variables/schema/retry_count), +并对章节级失败做有限重试;token 估算分块(真实拼回留待后续并发实现)。 +""" +from __future__ import annotations + +from genesis.inference.engine import InferenceEngine +from genesis.inference.prompt_registry import PromptRegistry +from genesis.inference.types import Prompt, StructuredResult +from genesis.writer.models import ChapterContent, GenerationContext +from genesis.writer.writer_state import WriterState +from genesis.writer.exceptions import WriterGenerationError + + +WRITER_PROMPT_TEMPLATE = ( + "你是概要设计书撰写专家。\n" + "章节: {{chapter_id}} {{title}}\n" + "写入规则:\n{{write_rules}}\n" + "设计规则:\n{{design_rules}}\n" + "模板样式:\n{{template_styles}}\n" + "参考资料:\n{{source}}\n" + "请输出符合 schema 的章节内容 JSON。" +) + +CHAPTER_OUTPUT_SCHEMA = { + "type": "object", + "properties": { + "title": {"type": "string"}, + "blocks": { + "type": "array", + "items": { + "type": "object", + "properties": { + "type": {"type": "string"}, + "text": {"type": "string"}, + "level": {"type": "integer"}, + "headers": {"type": "array", "items": {"type": "string"}}, + "rows": {"type": "array", "items": {"type": "array", "items": {"type": "string"}}}, + "items": {"type": "array", "items": {"type": "string"}}, + "source_uris": {"type": "array", "items": {"type": "string"}}, + }, + "required": ["type"], + }, + }, + }, + "required": ["title", "blocks"], +} + + +class WriterAgent: + def __init__( + self, + session_id, + engine: InferenceEngine, + prompt_registry: PromptRegistry, + state: WriterState, + max_retries: int = 1, + ) -> None: + self.session_id = session_id + self.engine = engine + self.prompt_registry = prompt_registry + self.state = state + self.max_retries = max_retries + + def _chunk_source(self, source) -> list[dict]: + if source is None: + return [{"index": 0, "text": ""}] + src = source if isinstance(source, str) else str(source) + n = max(1, len(src) // 1800 + 1) + return [{"index": i, "text": src[i * 1800:(i + 1) * 1800]} for i in range(n)] + + def _resolve_prompt(self) -> Prompt: + """取用/注册 writer.chapter 模板。 + + 优先使用 get_or_create(与测试 Fake 兼容);真实 PromptRegistry 无该方法时, + 回退为 register + get。 + """ + get_or_create = getattr(self.prompt_registry, "get_or_create", None) + if get_or_create is not None: + return get_or_create("writer.chapter", WRITER_PROMPT_TEMPLATE) + self.prompt_registry.register("writer.chapter", "1", WRITER_PROMPT_TEMPLATE) + return self.prompt_registry.get("writer.chapter", "1") + + def _call_llm(self, context: GenerationContext) -> dict: + prompt = self._resolve_prompt() + result: StructuredResult = self.engine.chat_structured( + session_id=self.session_id, + prompt=prompt, + variables=context.to_vars(), + schema=CHAPTER_OUTPUT_SCHEMA, + retry_count=2, + ) + if result.status not in ("ok", "fallback"): + raise WriterGenerationError(f"引擎返回异常状态: {result.status}") + return result.data + + def generate_chapter(self, context: GenerationContext) -> ChapterContent: + self._chunk_source(context.structured_source) # 分块可用性验证(真实拼回留待后续) + last_err: Exception | None = None + for _ in range(max(1, self.max_retries)): + try: + data = self._call_llm(context) + except Exception as e: # 引擎可能抛出任意异常,统一按章节级失败重试 + last_err = e + continue + try: + content = ChapterContent.from_llm(context.chapter_id, context.title, data) + except (KeyError, TypeError, ValueError) as e: + last_err = e + continue + self.state.record_success(content) + return content + raise WriterGenerationError(f"章节 {context.chapter_id} 重试耗尽: {last_err}")