fix(writer): await async InferenceEngine.chat_structured in WriterAgent
This commit is contained in:
@@ -5,6 +5,8 @@
|
|||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
from genesis.inference.engine import InferenceEngine
|
from genesis.inference.engine import InferenceEngine
|
||||||
from genesis.inference.prompt_registry import PromptRegistry
|
from genesis.inference.prompt_registry import PromptRegistry
|
||||||
from genesis.inference.types import Prompt, StructuredResult
|
from genesis.inference.types import Prompt, StructuredResult
|
||||||
@@ -84,13 +86,17 @@ class WriterAgent:
|
|||||||
|
|
||||||
def _call_llm(self, context: GenerationContext) -> dict:
|
def _call_llm(self, context: GenerationContext) -> dict:
|
||||||
prompt = self._resolve_prompt()
|
prompt = self._resolve_prompt()
|
||||||
result: StructuredResult = self.engine.chat_structured(
|
result = self.engine.chat_structured(
|
||||||
session_id=self.session_id,
|
session_id=self.session_id,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
variables=context.to_vars(),
|
variables=context.to_vars(),
|
||||||
schema=CHAPTER_OUTPUT_SCHEMA,
|
schema=CHAPTER_OUTPUT_SCHEMA,
|
||||||
retry_count=2,
|
retry_count=2,
|
||||||
)
|
)
|
||||||
|
# 真实 InferenceEngine.chat_structured 为 async;测试用同步 FakeEngine 返回普通对象。
|
||||||
|
# 兼容两者:若返回协程则通过 asyncio.run 驱动(调用方 orchestrator/qa_loop/脚本均为同步上下文)。
|
||||||
|
if asyncio.iscoroutine(result):
|
||||||
|
result = asyncio.run(result)
|
||||||
if result.status not in ("ok", "fallback"):
|
if result.status not in ("ok", "fallback"):
|
||||||
raise WriterGenerationError(f"引擎返回异常状态: {result.status}")
|
raise WriterGenerationError(f"引擎返回异常状态: {result.status}")
|
||||||
return result.data
|
return result.data
|
||||||
|
|||||||
Reference in New Issue
Block a user