fix: 收紧诊断 trace 输入边界

This commit is contained in:
root
2026-07-11 23:58:15 +02:00
parent 8fa0f093ec
commit 36dbfcab4a
3 changed files with 85 additions and 18 deletions
+10 -4
View File
@@ -208,13 +208,19 @@ function requiredSingleHeader(headers: AgentRunOtelRequestHeaders, name: string)
}
function singleHeader(headers: AgentRunOtelRequestHeaders, name: string): string | null {
const value = Object.entries(headers).find(([key]) => key.toLowerCase() === name)?.[1];
if (value === undefined) return null;
const entry = Object.entries(headers).find(([key]) => key.toLowerCase() === name);
if (entry === undefined) return null;
const value = entry[1];
if (value === undefined) throw diagnosticContextError(`${name} must not be empty`);
if (Array.isArray(value)) {
if (value.length !== 1) throw diagnosticContextError(`${name} must have exactly one value`);
return value[0]?.trim() || null;
const normalized = value[0]?.trim() ?? "";
if (normalized.length === 0) throw diagnosticContextError(`${name} must not be empty`);
return normalized;
}
return value.trim() || null;
const normalized = value.trim();
if (normalized.length === 0) throw diagnosticContextError(`${name} must not be empty`);
return normalized;
}
function diagnosticContextError(message: string): AgentRunError {
+9 -3
View File
@@ -917,9 +917,15 @@ async function route({ method, url, body, store, sourceCommit, authSummary, diag
const startedAt = Date.now();
const runId = runResultMatch[1] ?? "";
const run = await store.getRun(runId);
const commandId = url.searchParams.get("commandId");
if (diagnosticTraceContext !== null && commandId === null) {
throw diagnosticTargetMismatch("diagnostic run result requires an explicit commandId query");
const commandIds = url.searchParams.getAll("commandId");
const commandId = commandIds[0] ?? null;
if (diagnosticTraceContext !== null) {
if (commandIds.length !== 1 || commandId === null || commandId.length === 0) {
throw diagnosticTargetMismatch("diagnostic run result requires exactly one non-empty commandId query");
}
if (commandId !== diagnosticTraceContext.targetCommandId) {
throw diagnosticTargetMismatch("diagnostic run result commandId does not match the diagnostic target command");
}
}
const command = commandId ? await store.getCommand(commandId) : null;
const traceContext = await validatedDiagnosticTraceContext({ store, context: diagnosticTraceContext, run, command });
@@ -20,12 +20,21 @@ const DIAGNOSTIC_TRACE_B = "22222222222222222222222222222222";
const DIAGNOSTIC_PARENT_A = "aaaaaaaaaaaaaaaa";
const DIAGNOSTIC_PARENT_B = "bbbbbbbbbbbbbbbb";
class TrackingMemoryAgentRunStore extends MemoryAgentRunStore {
createRunCalls = 0;
override createRun(input: Parameters<MemoryAgentRunStore["createRun"]>[0]) {
this.createRunCalls += 1;
return super.createRun(input);
}
}
const selfTest: SelfTestCase = async () => {
const collector = await startOtlpCollector();
const previousEndpoint = process.env.AGENTRUN_OTEL_EXPORTER_OTLP_TRACES_ENDPOINT;
process.env.AGENTRUN_OTEL_EXPORTER_OTLP_TRACES_ENDPOINT = collector.endpoint;
const store = new MemoryAgentRunStore();
const run = store.createRun({
const store = new TrackingMemoryAgentRunStore();
const runInput: Parameters<MemoryAgentRunStore["createRun"]>[0] = {
tenantId: "unidesk",
projectId: "pikasTech/agentrun",
workspaceRef: { kind: "host-path", path: "/tmp/agentrun-diagnostic-selftest" },
@@ -39,7 +48,8 @@ const selfTest: SelfTestCase = async () => {
secretScope: { allowCredentialEcho: false, providerCredentials: [] },
},
traceSink: { kind: "hwlab", traceId: BUSINESS_TRACE_ID },
});
};
const run = store.createRun(runInput);
const command = store.createCommand(run.id, {
type: "turn",
payload: { prompt: "diagnostic trace context", traceId: BUSINESS_TRACE_ID },
@@ -69,6 +79,8 @@ const selfTest: SelfTestCase = async () => {
`/api/v1/runs/${run.id}/result?commandId=${command.id}`,
`/api/v1/runs/${run.id}/events?afterSeq=0&limit=100`,
];
const ordinaryCommandDetail = asRecord(await getData(manager.baseUrl, `/api/v1/runs/${run.id}/commands/${command.id}`));
assert.equal(ordinaryCommandDetail.id, command.id);
const ordinary = await Promise.all(paths.map((path) => getData(manager.baseUrl, path)));
await waitForSpanCount(collector.payloads, 3);
const businessContext = agentRunOtelTraceContext(run, command);
@@ -90,29 +102,54 @@ const selfTest: SelfTestCase = async () => {
assertDiagnosticSpans(spansForTrace(collector.payloads, DIAGNOSTIC_TRACE_A), DIAGNOSTIC_PARENT_A, businessContext, command.id);
assertDiagnosticSpans(spansForTrace(collector.payloads, DIAGNOSTIC_TRACE_B), DIAGNOSTIC_PARENT_B, businessContext, command.id);
const latestCommand = store.createCommand(run.id, {
type: "turn",
payload: { prompt: "diagnostic latest-command guard", traceId: BUSINESS_TRACE_ID },
idempotencyKey: "diagnostic-latest-command-guard",
});
store.finishCommand(latestCommand.id, { terminalStatus: "failed", failureKind: "backend-timeout", failureMessage: "latest command must not replace the diagnostic target" });
const spanCountBeforePinnedRunResult = allSpans(collector.payloads).length;
const pinnedRunResult = asRecord(await getData(manager.baseUrl, `/api/v1/runs/${run.id}/result?commandId=${command.id}`, diagnosticA));
assert.equal(pinnedRunResult.commandId, command.id);
assert.equal(pinnedRunResult.terminalStatus, "completed");
assert.notEqual(pinnedRunResult.commandId, latestCommand.id);
await waitForSpanCount(collector.payloads, spanCountBeforePinnedRunResult + 1);
const spanCountBeforeRejectedReads = allSpans(collector.payloads).length;
const wrongBusiness = diagnosticHeaders(command.id, "33333333333333333333333333333333", "cccccccccccccccc", "trc_wrong_target");
const wrongCommand = diagnosticHeaders("cmd_wrong_target", "44444444444444444444444444444444", "dddddddddddddddd", BUSINESS_TRACE_ID);
const writeDiagnostic = diagnosticHeaders(command.id, "55555555555555555555555555555555", "eeeeeeeeeeeeeeee", BUSINESS_TRACE_ID);
const createRunCallsBeforeRejectedWrites = store.createRunCalls;
await assertRejected(manager.baseUrl, paths[0] ?? "", wrongBusiness, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, paths[0] ?? "", wrongCommand, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}`, diagnosticA, "diagnostic-trace-route-unsupported");
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}/commands/${command.id}`, diagnosticA, "diagnostic-trace-route-unsupported");
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}/result`, diagnosticA, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, "/api/v1/runs", writeDiagnostic, "diagnostic-trace-context-write-denied", "POST");
await waitForSpanCount(collector.payloads, spanCountBeforeRejectedReads + 5);
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}/result?commandId=`, diagnosticA, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}/result?commandId=${command.id}&commandId=${command.id}`, diagnosticA, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, `/api/v1/runs/${run.id}/result?commandId=${latestCommand.id}`, diagnosticA, "diagnostic-trace-target-mismatch");
await assertRejected(manager.baseUrl, "/api/v1/runs", writeDiagnostic, "diagnostic-trace-context-write-denied", "POST", runInput);
await waitForSpanCount(collector.payloads, spanCountBeforeRejectedReads + 8);
assert.equal(spansForTrace(collector.payloads, businessContext.traceId).length, 3);
const rejectedSpans = allSpans(collector.payloads).filter((span) => span.name === "diagnostic_request_rejected");
assert.equal(rejectedSpans.length, 5);
assert.equal(rejectedSpans.length, 8);
for (const span of rejectedSpans) {
assert.equal((asRecord(span.status).code), 2);
assert.equal(Array.isArray(span.links), true);
assert.equal(spanAttributes(span).failureKind, "schema-invalid");
assert.match(String(spanAttributes(span)["diagnostic.rejection_reason"]), /^diagnostic-trace-/u);
}
const emptyDiagnostic = diagnosticHeadersWithValue(diagnosticA, "");
const whitespaceDiagnostic = diagnosticHeadersWithValue(diagnosticA, " ");
const repeatedDiagnostic = new Headers(diagnosticA);
repeatedDiagnostic.append(AGENTRUN_TRACE_PURPOSE_HEADER, "diagnostic");
await assertRejected(manager.baseUrl, "/api/v1/runs", emptyDiagnostic, "diagnostic-trace-context-invalid", "POST", runInput);
await assertRejected(manager.baseUrl, "/api/v1/runs", whitespaceDiagnostic, "diagnostic-trace-context-invalid", "POST", runInput);
await assertRejected(manager.baseUrl, "/api/v1/runs", repeatedDiagnostic, "diagnostic-trace-context-invalid", "POST", runInput);
assert.equal(store.createRunCalls, createRunCallsBeforeRejectedWrites);
const malformedTraceparent = { ...diagnosticA, traceparent: "00-00000000000000000000000000000000-0000000000000000-01" };
await assertRejected(manager.baseUrl, paths[0] ?? "", malformedTraceparent, "diagnostic-trace-context-invalid");
await new Promise((resolve) => setTimeout(resolve, 30));
assert.equal(allSpans(collector.payloads).length, spanCountBeforeRejectedReads + 5);
assert.equal(allSpans(collector.payloads).length, spanCountBeforeRejectedReads + 8);
return {
name: "diagnostic-trace-context",
@@ -123,8 +160,10 @@ const selfTest: SelfTestCase = async () => {
"diagnostic-result-events-authority-stable",
"diagnostic-target-mismatch-fail-closed",
"diagnostic-target-command-mismatch-fail-closed",
"diagnostic-unsupported-route-fail-closed",
"diagnostic-command-detail-remains-ordinary",
"diagnostic-run-result-command-query-fail-closed",
"diagnostic-write-request-fail-closed",
"diagnostic-empty-and-repeated-headers-fail-closed",
"diagnostic-rejection-error-span",
"malformed-diagnostic-context-no-span",
],
@@ -147,6 +186,16 @@ function diagnosticHeaders(commandId: string, traceId: string, parentSpanId: str
};
}
function diagnosticHeadersWithValue(base: Record<string, string>, value: string): Record<string, string> {
return {
traceparent: base.traceparent ?? "",
[AGENTRUN_TRACE_PURPOSE_HEADER]: value,
[AGENTRUN_DIAGNOSTIC_OPERATION_HEADER]: value,
[AGENTRUN_DIAGNOSTIC_TARGET_BUSINESS_TRACE_ID_HEADER]: value,
[AGENTRUN_DIAGNOSTIC_TARGET_COMMAND_ID_HEADER]: value,
};
}
async function getData(baseUrl: string, path: string, headers: Record<string, string> = {}): Promise<JsonValue> {
const response = await fetch(new URL(path, baseUrl), { headers });
const envelope = await response.json() as JsonRecord;
@@ -155,8 +204,14 @@ async function getData(baseUrl: string, path: string, headers: Record<string, st
return envelope.data ?? null;
}
async function assertRejected(baseUrl: string, path: string, headers: Record<string, string>, expectedReason: string, method = "GET"): Promise<void> {
const response = await fetch(new URL(path, baseUrl), { method, headers });
async function assertRejected(baseUrl: string, path: string, headers: Record<string, string> | Headers, expectedReason: string, method = "GET", body?: unknown): Promise<void> {
const requestHeaders = new Headers(headers);
if (body !== undefined) requestHeaders.set("content-type", "application/json");
const response = await fetch(new URL(path, baseUrl), {
method,
headers: requestHeaders,
...(body === undefined ? {} : { body: JSON.stringify(body) }),
});
const envelope = await response.json() as JsonRecord;
assert.equal(response.status, 400, JSON.stringify(envelope));
assert.equal(envelope.ok, false);