Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions apps/eval-harness/src/lib/generate-scenario.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -307,6 +307,7 @@ const concurrencyLayer = Layer.mergeAll(
}),
Layer.succeed(DiagramGenerationPolicy, {
concurrency: 1,
maxRepairAttempts: 1,
maxRetries: 0,
requestTimeoutMs: 1_000,
retryDelayMs: 0,
Expand Down
179 changes: 171 additions & 8 deletions packages/diagram/generation/src/lib/candidates.ts
Original file line number Diff line number Diff line change
Expand Up @@ -103,13 +103,28 @@ export function extractJsonObject(text: string): unknown {
return JSON.parse(text);
} catch {
const firstBrace = text.indexOf("{");
const lastBrace = text.lastIndexOf("}");

if (firstBrace === -1 || lastBrace === -1 || lastBrace <= firstBrace) {
if (firstBrace === -1) {
throw new Error("Model output did not contain a JSON object.");
}

return JSON.parse(text.slice(firstBrace, lastBrace + 1));
let depth = 0;
let escaped = false;
let inString = false;
for (let index = firstBrace; index < text.length; index += 1) {
const character = text[index];
if (inString) {
if (escaped) escaped = false;
else if (character === "\\") escaped = true;
else if (character === '"') inString = false;
continue;
}
if (character === '"') inString = true;
else if (character === "{") depth += 1;
else if (character === "}") {
depth -= 1;
if (depth === 0) return JSON.parse(text.slice(firstBrace, index + 1));
}
}
throw new Error("Model output did not contain one complete JSON object.");
}
}

Expand All @@ -131,6 +146,25 @@ function withSketchiDiagramStyle(input: unknown): unknown {
: input;
}

function normalizeGeneratedFlowchartInput(input: unknown): unknown {
const styled = withSketchiDiagramStyle(input);
if (!isUnknownRecord(styled) || !Array.isArray(styled["edges"])) {
return styled;
}
return {
...styled,
edges: styled["edges"].map((edge) =>
isUnknownRecord(edge) &&
typeof edge["label"] === "string" &&
edge["label"].trim().length === 0
? Object.fromEntries(
Object.entries(edge).filter(([key]) => key !== "label"),
)
: edge,
),
};
}

function firstString(values: readonly unknown[]): string | undefined {
return values.find((value): value is string => typeof value === "string");
}
Expand Down Expand Up @@ -160,7 +194,7 @@ export function responseErrorDiagnostic(raw: unknown): string | undefined {

export function parseGeneratedFlowchart(text: string): FlowchartDiagram {
return parseFlowchartDiagram(
withSketchiDiagramStyle(extractJsonObject(text)),
normalizeGeneratedFlowchartInput(extractJsonObject(text)),
);
}

Expand All @@ -175,7 +209,7 @@ export function parseGeneratedDiagram(
}
return parseMindmapDiagram(withSketchiDiagramStyle(extracted));
}
return parseFlowchartDiagram(withSketchiDiagramStyle(extracted));
return parseFlowchartDiagram(normalizeGeneratedFlowchartInput(extracted));
}

interface CandidateParseFailure {
Expand Down Expand Up @@ -281,7 +315,7 @@ function parseCandidateDiagram(text: string): CandidateParseResult {

const decoded = safeParseDiagramSchema(
FlowchartDiagramSchema,
withSketchiDiagramStyle(extracted),
normalizeGeneratedFlowchartInput(extracted),
);
if (!decoded.success) {
const diagnostics = decoded.error.issues.map(schemaIssueDiagnostic);
Expand Down Expand Up @@ -330,6 +364,135 @@ export function candidateFromText(
};
}

const EXPLICIT_MINIMUM_PATTERN =
/\bat least\s+(\d+|one|two|three|four|five|six|seven|eight|nine|ten)\s+(?:(?:distinct|labeled)\s+)?(steps?|nodes?|topics?|decisions?(?:\s+nodes?)?)\b/giu;
const LOOP_REQUIREMENT_PATTERN =
/\b(?:feedback|review|resubmission|retry|revision|remediation|investigation)\s+loop\b|\bloop(?:s|ed|ing)?\s+(?:back|through|to)\b|\breturns?\s+to\b/iu;
const NUMBER_WORD_COUNTS: Readonly<Record<string, number>> = {
eight: 8,
five: 5,
four: 4,
nine: 9,
one: 1,
seven: 7,
six: 6,
ten: 10,
three: 3,
two: 2,
};

export interface ExplicitRequestMinimum {
readonly expectedCount: number;
readonly expectedUnit: "decision nodes" | "nodes" | "topics";
readonly requestedCount: number;
readonly requestedUnit: string;
}

function requestedCount(value: string): number | undefined {
const numeric = Number.parseInt(value, 10);
return Number.isInteger(numeric)
? numeric
: NUMBER_WORD_COUNTS[value.toLowerCase()];
}

export function explicitRequestMinimums(
request: string,
diagramType: DiagramGenerationPrompt["type"],
): readonly ExplicitRequestMinimum[] {
return Array.from(request.matchAll(EXPLICIT_MINIMUM_PATTERN), (match) => {
const requested = match[1] ? requestedCount(match[1]) : undefined;
const unit = match[2]?.toLowerCase();
if (!requested || !unit) return undefined;

const decisionMinimum = /^decisions?(?:\s+nodes?)?$/.test(unit);
const applies =
(diagramType === "flowchart" &&
(/^(?:steps?|nodes?)$/.test(unit) || decisionMinimum)) ||
(diagramType === "mindmap" && /^topics?$/.test(unit));
if (!applies) return undefined;

const expectedUnit: ExplicitRequestMinimum["expectedUnit"] = decisionMinimum
? "decision nodes"
: diagramType === "flowchart"
? "nodes"
: "topics";
return {
expectedCount:
diagramType === "flowchart" && !decisionMinimum
? Math.min(requested, 24)
: requested,
expectedUnit,
requestedCount: requested,
requestedUnit: unit,
};
}).filter(
(minimum): minimum is ExplicitRequestMinimum => minimum !== undefined,
);
}

function flowchartHasDirectedCycle(diagram: FlowchartDiagram): boolean {
const adjacency = new Map<string, string[]>();
for (const edge of diagram.edges) {
adjacency.set(edge.source, [
...(adjacency.get(edge.source) ?? []),
edge.target,
]);
}
const pathExists = (start: string, destination: string): boolean => {
const pending = [start];
const visited = new Set<string>();
while (pending.length > 0) {
const current = pending.pop();
if (!current) continue;
if (current === destination) return true;
if (visited.has(current)) continue;
visited.add(current);
pending.push(...(adjacency.get(current) ?? []));
}
return false;
};
return diagram.edges.some((edge) => pathExists(edge.target, edge.source));
}

/** Turn a structurally valid result that misses explicit requirements into repair input. */
export function enforceCandidateRequestRequirements(
candidate: DiagramGenerationCandidate,
request: DiagramGenerationRequest,
): DiagramGenerationCandidate {
const diagram = candidate.diagram;
if (!diagram || diagram.type !== request.prompt.type) {
return candidate;
}
const requestDiagnostics = [
...explicitRequestMinimums(request.prompt.request, request.prompt.type)
.map((minimum) => {
const actual =
minimum.expectedUnit === "decision nodes"
? diagram.nodes.filter((node) => node.kind === "decision").length
: diagram.nodes.length;
if (actual >= minimum.expectedCount) return undefined;

return `request_minimum_not_met: requested at least ${minimum.requestedCount} ${minimum.requestedUnit}, but the generated ${diagram.type} contained ${actual}. Hint: return a complete diagram with at least ${minimum.expectedCount} ${minimum.expectedUnit}.`;
})
.filter((diagnostic): diagnostic is string => diagnostic !== undefined),
...(diagram.type === "flowchart" &&
LOOP_REQUIREMENT_PATTERN.test(request.prompt.request) &&
!flowchartHasDirectedCycle(diagram)
? [
"request_loop_not_met: prompt requires a retry or loop, but the generated flowchart contains no directed cycle. Hint: add a real back-edge from the loop path to the intended process or decision node, never the start node.",
]
: []),
];
if (requestDiagnostics.length === 0) return candidate;

const { diagram: _diagram, ...withoutDiagram } = candidate;
return {
...withoutDiagram,
diagnostics: [...candidate.diagnostics, ...requestDiagnostics],
error: "Generated diagram did not satisfy explicit request requirements.",
};
}

export function summarizeGenerationCandidate(
candidate: DiagramGenerationCandidate,
): DiagramGenerationCandidateSummary {
Expand Down
2 changes: 2 additions & 0 deletions packages/diagram/generation/src/lib/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ export class DiagramGenerationPolicyConfig extends Schema.Class<DiagramGeneratio
"DiagramGenerationPolicyConfig",
)({
concurrency: Schema.Number,
maxRepairAttempts: Schema.Number,
maxRetries: Schema.Number,
requestTimeoutMs: Schema.Number,
retryDelayMs: Schema.Number,
Expand All @@ -33,6 +34,7 @@ export class DiagramGenerationPolicy extends Context.Service<

export const diagramGenerationPolicyDefaults: DiagramGenerationPolicyConfig = {
concurrency: 2,
maxRepairAttempts: 2,
maxRetries: 2,
requestTimeoutMs: 30_000,
retryDelayMs: 250,
Expand Down
73 changes: 43 additions & 30 deletions packages/diagram/generation/src/lib/cloudflare-google-ai-studio.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import { Clock, Context, Effect, Layer, Schema } from "effect";

import {
candidateFromText,
enforceCandidateRequestRequirements,
responseErrorDiagnostic,
type DiagramGenerationCacheMode,
type DiagramGenerationRequest,
Expand Down Expand Up @@ -68,6 +69,7 @@ function cacheModeHeaders(
return cacheMode === "fresh"
? {
"Cache-Control": "no-store",
"cf-aig-skip-cache": "true",
Pragma: "no-cache",
}
: {};
Expand Down Expand Up @@ -186,20 +188,26 @@ const runGatewayAttempt = Effect.fn(
const usage = extractGeminiUsage(raw);
const finishReason = extractGeminiFinishReason(raw);

return candidateFromText({
diagnostics:
finishReason === "MAX_TOKENS"
? [
"output_truncated: Gemini stopped at the maximum output-token budget; regenerate the complete diagram.",
]
: [],
model,
provider: "cloudflare-google-ai-studio",
raw,
text,
cacheMode: request.cacheMode ?? "default",
...(usage ? { usage } : {}),
});
return enforceCandidateRequestRequirements(
candidateFromText({
diagnostics:
finishReason === "MAX_TOKENS"
? [
"output_truncated: Gemini stopped at the maximum output-token budget; regenerate the complete diagram.",
]
: [],
model,
provider: "cloudflare-google-ai-studio",
raw,
text,
...(finishReason === "MAX_TOKENS"
? { error: "Gemini output was truncated." }
: {}),
cacheMode: request.cacheMode ?? "default",
...(usage ? { usage } : {}),
}),
request,
);
});

export const CloudflareGoogleAiStudioClientLive = Layer.effect(
Expand All @@ -221,23 +229,28 @@ export const CloudflareGoogleAiStudioClientLive = Layer.effect(
"sketchi.scenario_id": request.prompt.id,
});
return yield* runDiagramGenerationWithPolicy(
Effect.gen(function* () {
const gateway = yield* Effect.try({
try: () => ai.gateway(config.gatewayId),
catch: (cause) =>
DiagramGenerationTransportError.make({
cause,
message: errorMessage(
(attemptRequest) =>
Effect.gen(function* () {
const gateway = yield* Effect.try({
try: () => ai.gateway(config.gatewayId),
catch: (cause) =>
DiagramGenerationTransportError.make({
cause,
"AI Gateway could not be initialized.",
),
operation: "ai.gateway",
provider: "cloudflare-google-ai-studio",
retryable: false,
}),
});
return runGatewayAttempt(gateway, config.collectLog, request);
}),
message: errorMessage(
cause,
"AI Gateway could not be initialized.",
),
operation: "ai.gateway",
provider: "cloudflare-google-ai-studio",
retryable: false,
}),
});
return runGatewayAttempt(
gateway,
config.collectLog,
attemptRequest,
);
}),
request,
"cloudflare-google-ai-studio",
policy,
Expand Down
Loading
Loading