llm.py 6.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163
  1. """文本 LLM:判断 / 拆分 / 解构用,走阿里云百炼 OpenAI-compatible /chat/completions。
  2. 只做一件事:给 system + user,返回解析好的 JSON dict。HTTP/鉴权风格同 extractor。
  3. """
  4. from __future__ import annotations
  5. import time
  6. from typing import Any, Callable, Optional
  7. import httpx
  8. from core.config import Settings
  9. from core.jsonio import extract_json_object
  10. from pipeline.tracing import TraceContext, TraceWriter, hash_prompt, redact_headers, timed_ms
  11. class LLMError(RuntimeError):
  12. pass
  13. def chat_json(
  14. system: str,
  15. user: str,
  16. *,
  17. model: Optional[str] = None,
  18. settings: Optional[Settings] = None,
  19. http_post: Callable[..., Any] = httpx.post,
  20. env_file: str = ".env",
  21. timeout: float = 60.0,
  22. trace_writer: TraceWriter | None = None,
  23. trace_context: TraceContext | None = None,
  24. trace_stage: str = "decode",
  25. trace_substage: str = "chat_json",
  26. prompt_name: str = "chat_json",
  27. ) -> dict:
  28. """调一次对话,强约束输出 JSON,返回解析后的 dict。带一次重试。"""
  29. settings = settings or Settings.from_env(env_file)
  30. api_key = settings.bailian_api_key
  31. if not api_key:
  32. raise LLMError("missing ALIYUN_BAILIAN_API_KEY")
  33. model = model or settings.llm_model
  34. messages = [
  35. {"role": "system", "content": system + (
  36. "\n只输出一个严格合法的 JSON 对象,不要解释或 markdown。"
  37. "字符串值要写在一行内,内部的换行写成 \\n、双引号写成 \\\",不要出现裸换行或裸双引号。"
  38. )},
  39. {"role": "user", "content": user},
  40. ]
  41. url = f"{settings.bailian_base_url.rstrip('/')}/chat/completions"
  42. headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
  43. last_exc: Optional[Exception] = None
  44. # 最多 3 轮:HTTP 错误重试;JSON 不合法则把坏输出喂回去做 self-repair
  45. for attempt in range(3):
  46. started = time.perf_counter()
  47. response_json: dict[str, Any] = {}
  48. request_payload = {
  49. "model": model,
  50. "messages": messages,
  51. "response_format": {"type": "json_object"},
  52. }
  53. try:
  54. resp = http_post(url, headers=headers, json=request_payload, timeout=timeout)
  55. resp.raise_for_status()
  56. response_json = resp.json()
  57. content = response_json["choices"][0]["message"]["content"]
  58. except httpx.HTTPError as exc:
  59. last_exc = exc
  60. if trace_writer is not None:
  61. trace_writer.llm_call(
  62. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  63. stage=trace_stage,
  64. substage=trace_substage,
  65. provider="bailian",
  66. model_name=model,
  67. endpoint=url,
  68. prompt_name=prompt_name,
  69. prompt_hash=hash_prompt(system),
  70. request_payload={"headers": redact_headers(headers), **request_payload},
  71. status="failed",
  72. error_message=str(exc),
  73. latency_ms=timed_ms(started),
  74. attempt_index=attempt + 1,
  75. )
  76. if attempt < 2:
  77. continue
  78. raise LLMError(f"llm_http_error: {exc}") from exc
  79. except (KeyError, IndexError, TypeError) as exc:
  80. if trace_writer is not None:
  81. trace_writer.llm_call(
  82. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  83. stage=trace_stage,
  84. substage=trace_substage,
  85. provider="bailian",
  86. model_name=model,
  87. endpoint=url,
  88. prompt_name=prompt_name,
  89. prompt_hash=hash_prompt(system),
  90. request_payload={"headers": redact_headers(headers), **request_payload},
  91. response_payload=response_json,
  92. status="failed",
  93. error_message=str(exc),
  94. latency_ms=timed_ms(started),
  95. attempt_index=attempt + 1,
  96. )
  97. raise LLMError(f"llm_response_invalid: {exc}") from exc
  98. try:
  99. parsed = extract_json_object(content)
  100. if trace_writer is not None:
  101. trace_writer.llm_call(
  102. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  103. stage=trace_stage,
  104. substage=trace_substage,
  105. provider="bailian",
  106. model_name=model,
  107. endpoint=url,
  108. prompt_name=prompt_name,
  109. prompt_hash=hash_prompt(system),
  110. request_payload={"headers": redact_headers(headers), **request_payload},
  111. response_payload=response_json,
  112. parsed_payload=parsed,
  113. status="done",
  114. latency_ms=timed_ms(started),
  115. attempt_index=attempt + 1,
  116. )
  117. return parsed
  118. except ValueError as exc:
  119. last_exc = exc
  120. if trace_writer is not None:
  121. trace_writer.llm_call(
  122. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  123. stage=trace_stage,
  124. substage=trace_substage,
  125. provider="bailian",
  126. model_name=model,
  127. endpoint=url,
  128. prompt_name=prompt_name,
  129. prompt_hash=hash_prompt(system),
  130. request_payload={"headers": redact_headers(headers), **request_payload},
  131. response_payload={"content": content},
  132. status="failed",
  133. error_message=str(exc),
  134. latency_ms=timed_ms(started),
  135. attempt_index=attempt + 1,
  136. )
  137. messages = messages + [
  138. {"role": "assistant", "content": content},
  139. {"role": "user", "content": (
  140. "上面的输出不是严格合法的 JSON。请只重新输出严格合法的 JSON,"
  141. "字符串内的换行写成 \\n、双引号写成 \\\",不要任何解释或 markdown。"
  142. )},
  143. ]
  144. raise LLMError(f"llm_json_unrepairable: {last_exc}")
  145. # 默认对话器类型:(system, user) -> dict。stage 可注入假实现做离线测试。
  146. ChatFn = Callable[[str, str], dict]
  147. def default_chat(env_file: str = ".env", model: Optional[str] = None) -> ChatFn:
  148. settings = Settings.from_env(env_file)
  149. return lambda system, user: chat_json(
  150. system, user, model=model, settings=settings
  151. )