348 lines
13 KiB
Python

"""DeepSeek Harness 会话 JSONL 到通用 FrameworkEvent 的确定性归一。"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Mapping
from framework.primitives.artifacts import payload_sha256, write_jsonl_atomic
class DshNormalizationError(ValueError):
"""DSH 会话工件不可安全解析或缺少必要身份。"""
@dataclass(frozen=True)
class DshModelCall:
"""DSH 会话中一条 assistant/message 的模型回合摘要。"""
actual_model_id: str
provider: str
usage: Mapping[str, Any] = field(default_factory=dict)
stop_reason: str | None = None
cost_usd: float | None = None
@dataclass
class DshToolCall:
"""DSH 会话中一条顶层 tool/call 的依赖摘要。"""
tool_call_id: str
name: str
args: Mapping[str, Any] | None = None
is_error: bool = False
@dataclass(frozen=True)
class DshSessionOutcome:
"""归一后的 DSH 会话事实,不包含 Muse 业务裁决。"""
session_id: str
final_text: str | None
model_calls: tuple[DshModelCall, ...]
tool_calls: tuple[DshToolCall, ...]
turns: int
unknown_event_types: tuple[str, ...]
event_count: int
artifact_path: str
artifact_sha256: str
def _mapping(value: Any) -> Mapping[str, Any]:
return value if isinstance(value, Mapping) else {}
def _text_from_content(content: Any) -> str:
if isinstance(content, str):
return content
if not isinstance(content, list):
return ""
parts: list[str] = []
for block in content:
if isinstance(block, Mapping) and block.get("type") == "text":
parts.append(str(block.get("text") or ""))
return "".join(parts)
def _route_from_header(data: Mapping[str, Any]) -> tuple[str, str] | None:
header = _mapping(data.get("header"))
config = _mapping(header.get("config"))
provider = str(config.get("provider") or "").strip()
model = str(config.get("model") or "").strip()
if provider and model:
return provider, model
return None
def _route_from_message(data: Mapping[str, Any]) -> tuple[str, str] | None:
message = _mapping(data.get("message"))
source = _mapping(message.get("source"))
provider = str(source.get("provider") or "").strip()
model = str(source.get("model") or "").strip()
if provider and model:
return provider, model
return None
def _usage_cost(usage: Mapping[str, Any]) -> float | None:
for value in (usage.get("costUsd"), usage.get("cost_usd")):
if value is None:
continue
try:
parsed = float(value)
except (TypeError, ValueError):
continue
if parsed >= 0:
return parsed
cost = usage.get("cost")
if isinstance(cost, Mapping):
try:
parsed = float(cost.get("total"))
except (TypeError, ValueError):
return None
return parsed if parsed >= 0 else None
return None
def _tool_result_is_error(data: Mapping[str, Any]) -> bool:
message = _mapping(data.get("message"))
for block in message.get("content") or []:
if not isinstance(block, Mapping):
continue
for nested in block.get("content") or []:
if isinstance(nested, Mapping) and bool(nested.get("isError")):
return True
if bool(block.get("isError")):
return True
return False
def _safe_details(raw: Mapping[str, Any]) -> dict[str, Any]:
"""只保留轨迹摘要,禁止把 prompt、response 或工具正文复制进通用事件。"""
event_type = str(raw.get("type") or "unknown")
data = _mapping(raw.get("data"))
details: dict[str, Any] = {"type": event_type}
for key in ("turn", "step", "callId", "name", "rootCallId", "parentCallId", "subCallId"):
if key in data and isinstance(data[key], (str, int, float, bool)):
details[key] = data[key]
if event_type == "request/context":
for key in ("provider", "model"):
if isinstance(data.get(key), str):
details[key] = data[key]
if event_type == "request/header":
route = _route_from_header(data)
if route is not None:
details.update({"provider": route[0], "model": route[1]})
if event_type == "turn/end":
reason = _mapping(data.get("reason"))
if isinstance(reason.get("kind"), str):
details["reason"] = reason["kind"]
if event_type == "tool/result":
details["isError"] = _tool_result_is_error(data)
if event_type == "tool/code-dispatch":
details["isError"] = bool(data.get("isError"))
return details
def _kind_phase(raw: Mapping[str, Any]) -> tuple[str, str]:
event_type = str(raw.get("type") or "unknown")
if event_type == "session":
return "session", "started"
if event_type in {"turn/start", "step/start"}:
return ("turn" if event_type.startswith("turn") else "step"), "started"
if event_type in {"turn/end", "step/end"}:
data = _mapping(raw.get("data"))
reason = _mapping(data.get("reason"))
failed = reason.get("kind") in {"error", "interrupted", "cancelled"}
return ("turn" if event_type.startswith("turn") else "step"), "failed" if failed else "completed"
if event_type == "assistant/message":
return "model", "completed"
if event_type == "assistant/chunk":
return "model", "progress"
if event_type == "tool/call" or event_type.endswith("/start") and event_type.startswith("tool/"):
return "tool", "started"
if (
event_type in {"tool/result", "tool/code-dispatch"}
or event_type.endswith("/end") and event_type.startswith("tool/")
):
data = _mapping(raw.get("data"))
failed = event_type == "tool/result" and _tool_result_is_error(data)
if event_type == "tool/code-dispatch":
failed = bool(data.get("isError"))
return "tool", "failed" if failed else "completed"
if event_type in {"user/message", "agent/inbox/spliced"} or event_type.startswith("agent/"):
return "agent", "progress"
if event_type.startswith(("request/", "permission/", "sandbox/", "approval/", "session/", "compaction/")):
return "transport", "progress"
if event_type in {"todo/write"} or event_type.startswith(("goal/", "workflow/", "tool-workflow/")):
return "transport", "progress"
return "unknown", "progress"
def _source_event_id(raw: Mapping[str, Any], line_number: int) -> str:
data = _mapping(raw.get("data"))
for candidate in (raw.get("id"), data.get("callId"), data.get("subCallId"), raw.get("seq")):
if candidate is not None and str(candidate).strip():
return str(candidate)
return f"line:{line_number}"
def _observed_at() -> str:
return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z")
def normalize_dsh_session(
session_path: str | Path,
artifact_path: str | Path,
*,
run_id: str | None = None,
framework_version: str = "unknown",
) -> DshSessionOutcome:
"""读取 DSH 明文 JSONL 会话,写出通用事件工件并返回执行摘要。"""
source = Path(session_path)
if source.suffix != ".jsonl":
raise DshNormalizationError(
"DSH_COMPRESSED_ARTIFACT_UNSUPPORTED: 适配器要求 compression=none 的明文会话"
)
try:
lines = source.read_text(encoding="utf-8").splitlines()
except (OSError, UnicodeError) as exc:
raise DshNormalizationError(f"DSH 会话不可读: {source}") from exc
if not lines:
raise DshNormalizationError("DSH_SESSION_EMPTY: 会话工件为空")
raw_events: list[dict[str, Any]] = []
for line_number, line in enumerate(lines, start=1):
if not line.strip():
continue
try:
raw = json.loads(line)
except json.JSONDecodeError as exc:
raise DshNormalizationError(
f"DSH_SESSION_INVALID: 第 {line_number} 行不是合法 JSON"
) from exc
if not isinstance(raw, Mapping):
raise DshNormalizationError(f"DSH_SESSION_INVALID: 第 {line_number} 行不是对象")
raw_events.append(dict(raw))
if not raw_events or raw_events[0].get("type") != "session":
raise DshNormalizationError("DSH_SESSION_INVALID: 首条记录不是 session header")
session_id = str(raw_events[0].get("id") or "").strip()
if not session_id:
raise DshNormalizationError("DSH_SESSION_INVALID: session header 缺少 id")
events: list[dict[str, Any]] = []
model_calls: list[DshModelCall] = []
tool_calls: list[DshToolCall] = []
pending_tools: dict[str, DshToolCall] = {}
current_route: tuple[str, str] | None = None
finish_reasons: dict[tuple[Any, Any], str] = {}
final_text: str | None = None
turns: set[int] = set()
unknown: set[str] = set()
for line_number, raw in enumerate(raw_events, start=1):
event_type = str(raw.get("type") or "unknown")
data = _mapping(raw.get("data"))
if event_type == "request/header":
current_route = _route_from_header(data) or current_route
elif event_type == "assistant/chunk":
chunk = _mapping(data.get("chunk"))
if chunk.get("type") == "finish":
reason = _mapping(chunk.get("reason"))
kind = reason.get("kind")
if isinstance(kind, str):
finish_reasons[(data.get("turn"), data.get("step"))] = kind
elif event_type == "assistant/message":
route = _route_from_message(data) or current_route
if route is None:
raise DshNormalizationError(
"DSH_MODEL_ID_MISSING: assistant/message 缺少 provider/model 路由"
)
message = _mapping(data.get("message"))
usage = _mapping(data.get("usage"))
model_calls.append(
DshModelCall(
actual_model_id=f"{route[0]}/{route[1]}",
provider=route[0],
usage=dict(usage),
stop_reason=finish_reasons.get((data.get("turn"), data.get("step"))),
cost_usd=_usage_cost(usage),
)
)
text = _text_from_content(message.get("content"))
if text:
final_text = text
elif event_type == "tool/call":
call_id = str(data.get("callId") or "").strip()
name = str(data.get("name") or "").strip()
if not call_id or not name:
raise DshNormalizationError("DSH_TOOL_EVENT_INVALID: tool/call 缺 callId/name")
args: Mapping[str, Any] | None = None
raw_args = data.get("arguments")
if isinstance(raw_args, str):
try:
parsed = json.loads(raw_args)
except json.JSONDecodeError:
parsed = None
if isinstance(parsed, Mapping):
args = dict(parsed)
call = DshToolCall(tool_call_id=call_id, name=name, args=args)
tool_calls.append(call)
pending_tools[call_id] = call
elif event_type == "tool/result":
message = _mapping(data.get("message"))
source = _mapping(message.get("source"))
call_id = str(source.get("callId") or "").strip()
if call_id in pending_tools:
pending_tools[call_id].is_error = _tool_result_is_error(data)
elif event_type == "turn/start":
turn = data.get("turn")
if isinstance(turn, int):
turns.add(turn)
kind, phase = _kind_phase(raw)
if kind == "unknown":
unknown.add(event_type)
events.append(
{
"framework": "dsh",
"frameworkVersion": framework_version,
"sessionId": session_id,
"sourceSeq": len(events) + 1,
"sourceEventId": _source_event_id(raw, line_number),
"runId": run_id,
"kind": kind,
"phase": phase,
"safeDetails": _safe_details(raw),
"payloadSha256": payload_sha256(raw),
"observedAt": _observed_at(),
}
)
artifact_sha256 = write_jsonl_atomic(artifact_path, events)
return DshSessionOutcome(
session_id=session_id,
final_text=final_text,
model_calls=tuple(model_calls),
tool_calls=tuple(tool_calls),
turns=len(turns),
unknown_event_types=tuple(sorted(unknown)),
event_count=len(events),
artifact_path=str(artifact_path),
artifact_sha256=artifact_sha256,
)
__all__ = [
"DshModelCall",
"DshNormalizationError",
"DshSessionOutcome",
"DshToolCall",
"normalize_dsh_session",
]