348 lines
13 KiB
Python
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",
|
|
]
|