mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
fix(knowledge-fs): improve semantic extraction resilience
This commit is contained in:
parent
5eabe14982
commit
bf5f916cd4
@ -0,0 +1,25 @@
|
||||
# Semantic Extraction Resilience
|
||||
|
||||
## Summary
|
||||
|
||||
- Bound entity and relation extraction to four concurrent model requests.
|
||||
- Retry malformed model JSON and retryable model-runtime failures with a bounded retry budget.
|
||||
- Convert LLM stream-read aborts into retryable Dify model-runtime timeout errors.
|
||||
- Validate all returned entities and relations, then retain the highest-confidence results within the configured per-node limit.
|
||||
|
||||
## Why
|
||||
|
||||
- A 52-node document previously launched enough concurrent model requests to exhaust Dify's database connection pool.
|
||||
- A stream crossing the request deadline during `reader.read()` leaked the raw abort `Symbol` instead of a retryable runtime error.
|
||||
- Model responses can occasionally contain malformed JSON or exceed the requested result count; neither condition should fail an otherwise valid document immediately.
|
||||
|
||||
## Verification
|
||||
|
||||
- Focused semantic extraction and Dify model-runtime client tests pass.
|
||||
- KnowledgeFS API and Dify model-runtime client typechecks pass.
|
||||
- Targeted Biome checks pass for the changed files.
|
||||
- The full KnowledgeFS check and build passed before separating the pre-existing failed-generation compatibility changes into a stash.
|
||||
|
||||
## Notes
|
||||
|
||||
- Compatibility logic for resuming pre-existing failed generation data is intentionally excluded from this change and stored in a separate git stash.
|
||||
70
knowledge-fs/packages/api/src/bounded-concurrency.test.ts
Normal file
70
knowledge-fs/packages/api/src/bounded-concurrency.test.ts
Normal file
@ -0,0 +1,70 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { createConcurrencyGate, mapWithConcurrency } from "./bounded-concurrency";
|
||||
|
||||
describe("bounded concurrency", () => {
|
||||
it("shares a concurrency gate across independent callers", async () => {
|
||||
let active = 0;
|
||||
let maxActive = 0;
|
||||
let release: (() => void) | undefined;
|
||||
const blocked = new Promise<void>((resolve) => {
|
||||
release = resolve;
|
||||
});
|
||||
let startedFour: (() => void) | undefined;
|
||||
const fourStarted = new Promise<void>((resolve) => {
|
||||
startedFour = resolve;
|
||||
});
|
||||
const gate = createConcurrencyGate(4);
|
||||
|
||||
const calls = Array.from({ length: 8 }, () =>
|
||||
gate.run(async () => {
|
||||
active += 1;
|
||||
maxActive = Math.max(maxActive, active);
|
||||
if (active === 4) {
|
||||
startedFour?.();
|
||||
}
|
||||
await blocked;
|
||||
active -= 1;
|
||||
}),
|
||||
);
|
||||
|
||||
await fourStarted;
|
||||
release?.();
|
||||
await Promise.all(calls);
|
||||
|
||||
expect(maxActive).toBe(4);
|
||||
});
|
||||
|
||||
it("stops scheduling and waits for active work after the first failure", async () => {
|
||||
const started: number[] = [];
|
||||
let release: (() => void) | undefined;
|
||||
const blocked = new Promise<void>((resolve) => {
|
||||
release = resolve;
|
||||
});
|
||||
let rejected = false;
|
||||
|
||||
const observed = mapWithConcurrency([0, 1, 2, 3, 4], 2, async (value) => {
|
||||
started.push(value);
|
||||
if (value === 0) {
|
||||
throw new Error("boom");
|
||||
}
|
||||
await blocked;
|
||||
return value;
|
||||
}).then(
|
||||
() => undefined,
|
||||
(error: unknown) => {
|
||||
rejected = true;
|
||||
return error;
|
||||
},
|
||||
);
|
||||
|
||||
await Promise.resolve();
|
||||
await Promise.resolve();
|
||||
expect(rejected).toBe(false);
|
||||
expect(started).toEqual([0, 1]);
|
||||
|
||||
release?.();
|
||||
await expect(observed).resolves.toMatchObject({ message: "boom" });
|
||||
expect(started).toEqual([0, 1]);
|
||||
});
|
||||
});
|
||||
85
knowledge-fs/packages/api/src/bounded-concurrency.ts
Normal file
85
knowledge-fs/packages/api/src/bounded-concurrency.ts
Normal file
@ -0,0 +1,85 @@
|
||||
export interface ConcurrencyGate {
|
||||
run<T>(fn: () => Promise<T>): Promise<T>;
|
||||
}
|
||||
|
||||
/** Fair FIFO gate that shares a fixed concurrency budget across independent callers. */
|
||||
export function createConcurrencyGate(limit: number): ConcurrencyGate {
|
||||
if (!Number.isSafeInteger(limit) || limit < 1) {
|
||||
throw new Error("Concurrency gate limit must be at least 1");
|
||||
}
|
||||
|
||||
let active = 0;
|
||||
const waiters: Array<() => void> = [];
|
||||
|
||||
const acquire = async (): Promise<void> => {
|
||||
if (active < limit) {
|
||||
active += 1;
|
||||
return;
|
||||
}
|
||||
|
||||
await new Promise<void>((resolve) => {
|
||||
waiters.push(resolve);
|
||||
});
|
||||
};
|
||||
|
||||
const release = (): void => {
|
||||
const next = waiters.shift();
|
||||
if (next) {
|
||||
next();
|
||||
return;
|
||||
}
|
||||
|
||||
active -= 1;
|
||||
};
|
||||
|
||||
return {
|
||||
run: async <T>(fn: () => Promise<T>): Promise<T> => {
|
||||
await acquire();
|
||||
try {
|
||||
return await fn();
|
||||
} finally {
|
||||
release();
|
||||
}
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
/** Order-preserving async map that stops scheduling after the first observed failure. */
|
||||
export async function mapWithConcurrency<T, R>(
|
||||
items: readonly T[],
|
||||
limit: number,
|
||||
fn: (item: T, index: number) => Promise<R>,
|
||||
): Promise<R[]> {
|
||||
const results = new Array<R>(items.length);
|
||||
let cursor = 0;
|
||||
let failed = false;
|
||||
let firstError: unknown;
|
||||
|
||||
async function worker(): Promise<void> {
|
||||
while (!failed) {
|
||||
const index = cursor;
|
||||
cursor += 1;
|
||||
if (index >= items.length) {
|
||||
return;
|
||||
}
|
||||
|
||||
try {
|
||||
results[index] = await fn(items[index] as T, index);
|
||||
} catch (error) {
|
||||
if (!failed) {
|
||||
failed = true;
|
||||
firstError = error;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const workerCount = Math.max(1, Math.min(limit, items.length));
|
||||
await Promise.all(Array.from({ length: workerCount }, () => worker()));
|
||||
|
||||
if (failed) {
|
||||
throw firstError;
|
||||
}
|
||||
|
||||
return results;
|
||||
}
|
||||
@ -945,7 +945,7 @@ describe("entity extraction", () => {
|
||||
});
|
||||
});
|
||||
|
||||
it("rejects invalid or unbounded entity extraction inputs and provider output", async () => {
|
||||
it("rejects invalid entity extraction inputs and bounds provider output", async () => {
|
||||
const nodes = createInMemoryKnowledgeNodeRepository({
|
||||
maxBatchSize: 2,
|
||||
maxListLimit: 2,
|
||||
@ -961,8 +961,8 @@ describe("entity extraction", () => {
|
||||
provider: {
|
||||
extract: async () => ({
|
||||
entities: [
|
||||
{ confidence: 0.9, text: "Refund Policy", type: "policy" },
|
||||
{ confidence: 0.8, text: "Acme", type: "organization" },
|
||||
{ confidence: 0.9, text: "Refund Policy", type: "policy" },
|
||||
],
|
||||
}),
|
||||
},
|
||||
@ -979,7 +979,21 @@ describe("entity extraction", () => {
|
||||
knowledgeSpaceId: first.knowledgeSpaceId,
|
||||
nodeIds: [first.id],
|
||||
}),
|
||||
).rejects.toThrow("Entity extraction provider returned 2 entities over maxEntitiesPerNode=1");
|
||||
).resolves.toMatchObject({
|
||||
extractedNodes: [
|
||||
{
|
||||
metadata: {
|
||||
extractedEntities: [
|
||||
{
|
||||
confidence: 0.9,
|
||||
text: "Refund Policy",
|
||||
type: "policy",
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
});
|
||||
await expect(
|
||||
createEntityExtractionFlow({
|
||||
maxBatchSize: 1,
|
||||
@ -1248,7 +1262,7 @@ describe("relation extraction", () => {
|
||||
});
|
||||
});
|
||||
|
||||
it("rejects invalid or unbounded relation extraction inputs and provider output", async () => {
|
||||
it("rejects invalid relation extraction inputs and bounds provider output", async () => {
|
||||
const nodes = createInMemoryKnowledgeNodeRepository({
|
||||
maxBatchSize: 2,
|
||||
maxListLimit: 2,
|
||||
@ -1264,8 +1278,8 @@ describe("relation extraction", () => {
|
||||
provider: {
|
||||
extract: async () => ({
|
||||
relations: [
|
||||
{ confidence: 0.9, object: "B", subject: "A", type: "mentions" },
|
||||
{ confidence: 0.8, object: "D", subject: "C", type: "references" },
|
||||
{ confidence: 0.9, object: "B", subject: "A", type: "mentions" },
|
||||
],
|
||||
}),
|
||||
},
|
||||
@ -1282,9 +1296,22 @@ describe("relation extraction", () => {
|
||||
knowledgeSpaceId: first.knowledgeSpaceId,
|
||||
nodeIds: [first.id],
|
||||
}),
|
||||
).rejects.toThrow(
|
||||
"Relation extraction provider returned 2 relations over maxRelationsPerNode=1",
|
||||
);
|
||||
).resolves.toMatchObject({
|
||||
extractedNodes: [
|
||||
{
|
||||
metadata: {
|
||||
extractedRelations: [
|
||||
{
|
||||
confidence: 0.9,
|
||||
object: "B",
|
||||
subject: "A",
|
||||
type: "mentions",
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
});
|
||||
await expect(
|
||||
createRelationExtractionFlow({
|
||||
maxBatchSize: 1,
|
||||
|
||||
@ -1,5 +1,6 @@
|
||||
import { type KnowledgeNode, PublicationGenerationIdSchema } from "@knowledge/core";
|
||||
|
||||
import { mapWithConcurrency } from "./bounded-concurrency";
|
||||
import { ENTITY_EXTRACTION_TYPES, type EntityExtractionType } from "./extraction-types";
|
||||
import { cloneJsonObject, isPlainObject } from "./json-utils";
|
||||
import { type KnowledgeNodeRepository, cloneKnowledgeNode } from "./knowledge-node-repository";
|
||||
@ -31,6 +32,7 @@ export interface EntityExtractionProvider {
|
||||
|
||||
export interface EntityExtractionFlowOptions {
|
||||
readonly maxBatchSize: number;
|
||||
readonly maxConcurrency?: number | undefined;
|
||||
readonly maxEntitiesPerNode?: number | undefined;
|
||||
readonly model: string;
|
||||
readonly nodes: KnowledgeNodeRepository;
|
||||
@ -58,6 +60,7 @@ export interface EntityExtractionFlow {
|
||||
|
||||
export function createEntityExtractionFlow({
|
||||
maxBatchSize,
|
||||
maxConcurrency = 4,
|
||||
maxEntitiesPerNode = 100,
|
||||
model,
|
||||
nodes,
|
||||
@ -69,6 +72,10 @@ export function createEntityExtractionFlow({
|
||||
throw new Error("Entity extraction maxBatchSize must be at least 1");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxConcurrency) || maxConcurrency < 1) {
|
||||
throw new Error("Entity extraction maxConcurrency must be at least 1");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxEntitiesPerNode) || maxEntitiesPerNode < 1) {
|
||||
throw new Error("Entity extraction maxEntitiesPerNode must be at least 1");
|
||||
}
|
||||
@ -110,32 +117,30 @@ export function createEntityExtractionFlow({
|
||||
};
|
||||
}
|
||||
|
||||
const generated = await Promise.all(
|
||||
orderedNodes.map(async (node) => {
|
||||
const result = await provider.extract({
|
||||
maxEntities: maxEntitiesPerNode,
|
||||
model,
|
||||
node: cloneKnowledgeNode(node),
|
||||
prompt: entityExtractionPrompt(node),
|
||||
promptVersion,
|
||||
...(tenantId ? { tenantId } : {}),
|
||||
});
|
||||
const entities = validateExtractedEntities(result.entities, maxEntitiesPerNode);
|
||||
const generated = await mapWithConcurrency(orderedNodes, maxConcurrency, async (node) => {
|
||||
const result = await provider.extract({
|
||||
maxEntities: maxEntitiesPerNode,
|
||||
model,
|
||||
node: cloneKnowledgeNode(node),
|
||||
prompt: entityExtractionPrompt(node),
|
||||
promptVersion,
|
||||
...(tenantId ? { tenantId } : {}),
|
||||
});
|
||||
const entities = validateExtractedEntities(result.entities, maxEntitiesPerNode);
|
||||
|
||||
return {
|
||||
id: node.id,
|
||||
metadata: entityExtractionMetadata({
|
||||
entities,
|
||||
metadata: result.metadata,
|
||||
model,
|
||||
node,
|
||||
now,
|
||||
promptVersion,
|
||||
traceId,
|
||||
}),
|
||||
};
|
||||
}),
|
||||
);
|
||||
return {
|
||||
id: node.id,
|
||||
metadata: entityExtractionMetadata({
|
||||
entities,
|
||||
metadata: result.metadata,
|
||||
model,
|
||||
node,
|
||||
now,
|
||||
promptVersion,
|
||||
traceId,
|
||||
}),
|
||||
};
|
||||
});
|
||||
const extractedNodes = await nodes.updateMetadataMany({
|
||||
knowledgeSpaceId,
|
||||
patches: generated,
|
||||
@ -188,13 +193,7 @@ function validateExtractedEntities(
|
||||
entities: readonly ExtractedEntity[],
|
||||
maxEntitiesPerNode: number,
|
||||
): ExtractedEntity[] {
|
||||
if (entities.length > maxEntitiesPerNode) {
|
||||
throw new Error(
|
||||
`Entity extraction provider returned ${entities.length} entities over maxEntitiesPerNode=${maxEntitiesPerNode}`,
|
||||
);
|
||||
}
|
||||
|
||||
return entities.map((entity) => {
|
||||
const validated = entities.map((entity, index) => {
|
||||
if (!ENTITY_EXTRACTION_TYPES.has(entity.type)) {
|
||||
throw new Error("Entity extraction entity type is unsupported");
|
||||
}
|
||||
@ -208,12 +207,22 @@ function validateExtractedEntities(
|
||||
}
|
||||
|
||||
return {
|
||||
confidence: entity.confidence,
|
||||
...(entity.metadata ? { metadata: cloneJsonObject(entity.metadata) } : {}),
|
||||
text: entity.text.trim(),
|
||||
type: entity.type,
|
||||
entity: {
|
||||
confidence: entity.confidence,
|
||||
...(entity.metadata ? { metadata: cloneJsonObject(entity.metadata) } : {}),
|
||||
text: entity.text.trim(),
|
||||
type: entity.type,
|
||||
},
|
||||
index,
|
||||
};
|
||||
});
|
||||
|
||||
return validated
|
||||
.sort(
|
||||
(left, right) => right.entity.confidence - left.entity.confidence || left.index - right.index,
|
||||
)
|
||||
.slice(0, maxEntitiesPerNode)
|
||||
.map(({ entity }) => entity);
|
||||
}
|
||||
|
||||
function entityExtractionMetadata({
|
||||
|
||||
@ -1,6 +1,7 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { createLlmEntityExtractionProvider } from "./llm-entity-extraction-provider";
|
||||
import { createLlmRelationExtractionProvider } from "./llm-relation-extraction-provider";
|
||||
|
||||
describe("createLlmEntityExtractionProvider", () => {
|
||||
it("adapts strict LLM JSON into entity extraction provider results", async () => {
|
||||
@ -81,4 +82,128 @@ describe("createLlmEntityExtractionProvider", () => {
|
||||
}),
|
||||
).rejects.toThrow("LLM entity extraction provider returned invalid entity JSON");
|
||||
});
|
||||
|
||||
it("retries malformed JSON with an explicit correction turn", async () => {
|
||||
const calls: unknown[] = [];
|
||||
const provider = createLlmEntityExtractionProvider({
|
||||
provider: {
|
||||
generate: async (input) => {
|
||||
calls.push(input);
|
||||
|
||||
return {
|
||||
text:
|
||||
calls.length === 1
|
||||
? '{"entities":[{"text":"Acme","type":"organization","confidence":0.9}'
|
||||
: '{"entities":[{"text":"Acme","type":"organization","confidence":0.9}]}',
|
||||
};
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
await expect(
|
||||
provider.extract({
|
||||
maxEntities: 5,
|
||||
model: "entity-llm",
|
||||
node: {} as never,
|
||||
prompt: "Text: Acme",
|
||||
promptVersion: "entity-extraction-v1",
|
||||
}),
|
||||
).resolves.toMatchObject({
|
||||
entities: [expect.objectContaining({ text: "Acme", type: "organization" })],
|
||||
});
|
||||
expect(calls).toHaveLength(2);
|
||||
expect(calls[1]).toMatchObject({
|
||||
messages: expect.arrayContaining([
|
||||
expect.objectContaining({
|
||||
content: expect.stringContaining("invalid"),
|
||||
role: "user",
|
||||
}),
|
||||
]),
|
||||
});
|
||||
});
|
||||
|
||||
it("retries a retryable model runtime failure", async () => {
|
||||
let calls = 0;
|
||||
const provider = createLlmEntityExtractionProvider({
|
||||
provider: {
|
||||
generate: async () => {
|
||||
calls += 1;
|
||||
if (calls === 1) {
|
||||
throw Object.assign(new Error("model request timed out"), { retryable: true });
|
||||
}
|
||||
|
||||
return { text: '{"entities":[]}' };
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
await expect(
|
||||
provider.extract({
|
||||
maxEntities: 5,
|
||||
model: "entity-llm",
|
||||
node: {} as never,
|
||||
prompt: "Text: Acme",
|
||||
promptVersion: "entity-extraction-v1",
|
||||
}),
|
||||
).resolves.toMatchObject({ entities: [] });
|
||||
expect(calls).toBe(2);
|
||||
});
|
||||
|
||||
it("shares one model request limit across entity and relation providers", async () => {
|
||||
let active = 0;
|
||||
let maxActive = 0;
|
||||
let release: (() => void) | undefined;
|
||||
const blocked = new Promise<void>((resolve) => {
|
||||
release = resolve;
|
||||
});
|
||||
let startedFour: (() => void) | undefined;
|
||||
const fourStarted = new Promise<void>((resolve) => {
|
||||
startedFour = resolve;
|
||||
});
|
||||
const generate = async (input: {
|
||||
readonly messages: readonly { readonly content: string }[];
|
||||
}) => {
|
||||
active += 1;
|
||||
maxActive = Math.max(maxActive, active);
|
||||
if (active === 4) {
|
||||
startedFour?.();
|
||||
}
|
||||
await blocked;
|
||||
active -= 1;
|
||||
|
||||
return {
|
||||
text: input.messages[0]?.content.includes("relations")
|
||||
? '{"relations":[]}'
|
||||
: '{"entities":[]}',
|
||||
};
|
||||
};
|
||||
const entityProvider = createLlmEntityExtractionProvider({ provider: { generate } });
|
||||
const relationProvider = createLlmRelationExtractionProvider({ provider: { generate } });
|
||||
|
||||
const entityCalls = Array.from({ length: 4 }, () =>
|
||||
entityProvider.extract({
|
||||
maxEntities: 5,
|
||||
model: "shared-llm",
|
||||
node: {} as never,
|
||||
prompt: "Text: Acme",
|
||||
promptVersion: "entity-extraction-v1",
|
||||
}),
|
||||
);
|
||||
const relationCalls = Array.from({ length: 4 }, () =>
|
||||
relationProvider.extract({
|
||||
entities: [],
|
||||
maxRelations: 5,
|
||||
model: "shared-llm",
|
||||
node: {} as never,
|
||||
prompt: "Text: Acme",
|
||||
promptVersion: "relation-extraction-v1",
|
||||
}),
|
||||
);
|
||||
|
||||
await fourStarted;
|
||||
release?.();
|
||||
await Promise.all([...entityCalls, ...relationCalls]);
|
||||
|
||||
expect(maxActive).toBe(4);
|
||||
});
|
||||
});
|
||||
|
||||
@ -5,6 +5,7 @@ import type {
|
||||
EntityExtractionProviderInput,
|
||||
} from "./entity-extraction-flow";
|
||||
import { cloneJsonObject, isPlainObject } from "./json-utils";
|
||||
import { semanticExtractionModelRequestGate } from "./semantic-extraction-concurrency";
|
||||
|
||||
export interface LlmEntityExtractionMessage {
|
||||
readonly content: string;
|
||||
@ -33,15 +34,21 @@ export interface EntityExtractionTextProvider {
|
||||
|
||||
export interface LlmEntityExtractionProviderOptions {
|
||||
readonly maxOutputTokens?: number | undefined;
|
||||
readonly maxRetries?: number | undefined;
|
||||
readonly provider: EntityExtractionTextProvider;
|
||||
readonly temperature?: number | undefined;
|
||||
}
|
||||
|
||||
export function createLlmEntityExtractionProvider({
|
||||
maxOutputTokens = 1_500,
|
||||
maxRetries = 2,
|
||||
provider,
|
||||
temperature = 0,
|
||||
}: LlmEntityExtractionProviderOptions): EntityExtractionProvider {
|
||||
if (!Number.isInteger(maxRetries) || maxRetries < 0) {
|
||||
throw new Error("LLM entity extraction maxRetries must be a non-negative integer");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxOutputTokens) || maxOutputTokens < 1) {
|
||||
throw new Error("LLM entity extraction maxOutputTokens must be at least 1");
|
||||
}
|
||||
@ -52,14 +59,37 @@ export function createLlmEntityExtractionProvider({
|
||||
|
||||
return {
|
||||
extract: async (input) => {
|
||||
const result = await provider.generate({
|
||||
maxOutputTokens,
|
||||
messages: entityExtractionMessages(input),
|
||||
model: input.model,
|
||||
temperature,
|
||||
...(input.tenantId ? { tenantId: input.tenantId } : {}),
|
||||
});
|
||||
const parsed = parseLlmEntityExtractionJson(result.text);
|
||||
let messages = entityExtractionMessages(input);
|
||||
let result: GenerateEntityExtractionTextResult | undefined;
|
||||
let parsed: LlmEntityExtractionOutput | undefined;
|
||||
for (let attempt = 0; attempt <= maxRetries; attempt += 1) {
|
||||
try {
|
||||
result = await semanticExtractionModelRequestGate.run(() =>
|
||||
provider.generate({
|
||||
maxOutputTokens,
|
||||
messages,
|
||||
model: input.model,
|
||||
temperature,
|
||||
...(input.tenantId ? { tenantId: input.tenantId } : {}),
|
||||
}),
|
||||
);
|
||||
parsed = parseLlmEntityExtractionJson(result.text);
|
||||
break;
|
||||
} catch (error) {
|
||||
if (attempt >= maxRetries) {
|
||||
throw error;
|
||||
}
|
||||
if (result) {
|
||||
messages = entityExtractionCorrectionMessages(messages, result.text);
|
||||
result = undefined;
|
||||
} else if (!isRetryableModelError(error)) {
|
||||
throw error;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!result || !parsed) {
|
||||
throw new Error("LLM entity extraction format retry did not produce a result");
|
||||
}
|
||||
|
||||
return {
|
||||
entities: parsed.entities.map((entity) => ({
|
||||
@ -107,10 +137,30 @@ function entityExtractionMessages(
|
||||
];
|
||||
}
|
||||
|
||||
function parseLlmEntityExtractionJson(text: string): LlmEntityExtractionOutput {
|
||||
const parsed = tryParseJsonObject(text);
|
||||
function isRetryableModelError(error: unknown): boolean {
|
||||
return (
|
||||
typeof error === "object" && error !== null && "retryable" in error && error.retryable === true
|
||||
);
|
||||
}
|
||||
|
||||
function entityExtractionCorrectionMessages(
|
||||
messages: readonly LlmEntityExtractionMessage[],
|
||||
invalidText: string,
|
||||
): readonly LlmEntityExtractionMessage[] {
|
||||
return [
|
||||
...messages,
|
||||
{ content: invalidText.slice(0, 8_000), role: "assistant" },
|
||||
{
|
||||
content:
|
||||
"The previous response is invalid JSON or does not match the required schema. Return a corrected complete JSON object only.",
|
||||
role: "user",
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
function parseLlmEntityExtractionJson(text: string): LlmEntityExtractionOutput {
|
||||
try {
|
||||
const parsed = tryParseJsonObject(text);
|
||||
return LlmEntityExtractionOutputSchema.parse(parsed);
|
||||
} catch (error) {
|
||||
throw new Error("LLM entity extraction provider returned invalid entity JSON", {
|
||||
|
||||
@ -0,0 +1,45 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { createLlmRelationExtractionProvider } from "./llm-relation-extraction-provider";
|
||||
|
||||
describe("createLlmRelationExtractionProvider", () => {
|
||||
it("retries malformed JSON with an explicit correction turn", async () => {
|
||||
const calls: unknown[] = [];
|
||||
const provider = createLlmRelationExtractionProvider({
|
||||
provider: {
|
||||
generate: async (input) => {
|
||||
calls.push(input);
|
||||
|
||||
return {
|
||||
text:
|
||||
calls.length === 1
|
||||
? '{"relations":[{"subject":"Acme","type":"mentions","object":"React","confidence":0.9}'
|
||||
: '{"relations":[{"subject":"Acme","type":"mentions","object":"React","confidence":0.9}]}',
|
||||
};
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
await expect(
|
||||
provider.extract({
|
||||
entities: [],
|
||||
maxRelations: 5,
|
||||
model: "relation-llm",
|
||||
node: {} as never,
|
||||
prompt: "Text: Acme mentions React",
|
||||
promptVersion: "relation-extraction-v1",
|
||||
}),
|
||||
).resolves.toMatchObject({
|
||||
relations: [expect.objectContaining({ object: "React", subject: "Acme", type: "mentions" })],
|
||||
});
|
||||
expect(calls).toHaveLength(2);
|
||||
expect(calls[1]).toMatchObject({
|
||||
messages: expect.arrayContaining([
|
||||
expect.objectContaining({
|
||||
content: expect.stringContaining("invalid"),
|
||||
role: "user",
|
||||
}),
|
||||
]),
|
||||
});
|
||||
});
|
||||
});
|
||||
@ -4,6 +4,7 @@ import type {
|
||||
RelationExtractionProvider,
|
||||
RelationExtractionProviderInput,
|
||||
} from "./relation-extraction-flow";
|
||||
import { semanticExtractionModelRequestGate } from "./semantic-extraction-concurrency";
|
||||
|
||||
export interface LlmRelationExtractionMessage {
|
||||
readonly content: string;
|
||||
@ -34,15 +35,21 @@ export interface RelationExtractionTextProvider {
|
||||
|
||||
export interface LlmRelationExtractionProviderOptions {
|
||||
readonly maxOutputTokens?: number | undefined;
|
||||
readonly maxRetries?: number | undefined;
|
||||
readonly provider: RelationExtractionTextProvider;
|
||||
readonly temperature?: number | undefined;
|
||||
}
|
||||
|
||||
export function createLlmRelationExtractionProvider({
|
||||
maxOutputTokens = 1_500,
|
||||
maxRetries = 2,
|
||||
provider,
|
||||
temperature = 0,
|
||||
}: LlmRelationExtractionProviderOptions): RelationExtractionProvider {
|
||||
if (!Number.isInteger(maxRetries) || maxRetries < 0) {
|
||||
throw new Error("LLM relation extraction maxRetries must be a non-negative integer");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxOutputTokens) || maxOutputTokens < 1) {
|
||||
throw new Error("LLM relation extraction maxOutputTokens must be at least 1");
|
||||
}
|
||||
@ -53,14 +60,37 @@ export function createLlmRelationExtractionProvider({
|
||||
|
||||
return {
|
||||
extract: async (input) => {
|
||||
const result = await provider.generate({
|
||||
maxOutputTokens,
|
||||
messages: relationExtractionMessages(input),
|
||||
model: input.model,
|
||||
temperature,
|
||||
...(input.tenantId ? { tenantId: input.tenantId } : {}),
|
||||
});
|
||||
const parsed = parseLlmRelationExtractionJson(result.text);
|
||||
let messages = relationExtractionMessages(input);
|
||||
let result: GenerateRelationExtractionTextResult | undefined;
|
||||
let parsed: LlmRelationExtractionOutput | undefined;
|
||||
for (let attempt = 0; attempt <= maxRetries; attempt += 1) {
|
||||
try {
|
||||
result = await semanticExtractionModelRequestGate.run(() =>
|
||||
provider.generate({
|
||||
maxOutputTokens,
|
||||
messages,
|
||||
model: input.model,
|
||||
temperature,
|
||||
...(input.tenantId ? { tenantId: input.tenantId } : {}),
|
||||
}),
|
||||
);
|
||||
parsed = parseLlmRelationExtractionJson(result.text);
|
||||
break;
|
||||
} catch (error) {
|
||||
if (attempt >= maxRetries) {
|
||||
throw error;
|
||||
}
|
||||
if (result) {
|
||||
messages = relationExtractionCorrectionMessages(messages, result.text);
|
||||
result = undefined;
|
||||
} else if (!isRetryableModelError(error)) {
|
||||
throw error;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!result || !parsed) {
|
||||
throw new Error("LLM relation extraction format retry did not produce a result");
|
||||
}
|
||||
|
||||
return {
|
||||
metadata: {
|
||||
@ -104,10 +134,30 @@ function relationExtractionMessages(
|
||||
];
|
||||
}
|
||||
|
||||
function parseLlmRelationExtractionJson(text: string): LlmRelationExtractionOutput {
|
||||
const parsed = tryParseJsonObject(text);
|
||||
function isRetryableModelError(error: unknown): boolean {
|
||||
return (
|
||||
typeof error === "object" && error !== null && "retryable" in error && error.retryable === true
|
||||
);
|
||||
}
|
||||
|
||||
function relationExtractionCorrectionMessages(
|
||||
messages: readonly LlmRelationExtractionMessage[],
|
||||
invalidText: string,
|
||||
): readonly LlmRelationExtractionMessage[] {
|
||||
return [
|
||||
...messages,
|
||||
{ content: invalidText.slice(0, 8_000), role: "assistant" },
|
||||
{
|
||||
content:
|
||||
"The previous response is invalid JSON or does not match the required schema. Return a corrected complete JSON object only.",
|
||||
role: "user",
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
function parseLlmRelationExtractionJson(text: string): LlmRelationExtractionOutput {
|
||||
try {
|
||||
const parsed = tryParseJsonObject(text);
|
||||
return LlmRelationExtractionOutputSchema.parse(parsed);
|
||||
} catch (error) {
|
||||
throw new Error("LLM relation extraction provider returned invalid relation JSON", {
|
||||
|
||||
@ -1,5 +1,6 @@
|
||||
import { type KnowledgeNode, PublicationGenerationIdSchema } from "@knowledge/core";
|
||||
|
||||
import { mapWithConcurrency } from "./bounded-concurrency";
|
||||
import { type ExtractedEntity, extractedEntitiesFromNodeMetadata } from "./entity-extraction-flow";
|
||||
import { RELATION_EXTRACTION_TYPES, type RelationExtractionType } from "./extraction-types";
|
||||
import { cloneJsonObject, isPlainObject } from "./json-utils";
|
||||
@ -34,6 +35,7 @@ export interface RelationExtractionProvider {
|
||||
|
||||
export interface RelationExtractionFlowOptions {
|
||||
readonly maxBatchSize: number;
|
||||
readonly maxConcurrency?: number | undefined;
|
||||
readonly maxRelationsPerNode?: number | undefined;
|
||||
readonly model: string;
|
||||
readonly nodes: KnowledgeNodeRepository;
|
||||
@ -61,6 +63,7 @@ export interface RelationExtractionFlow {
|
||||
|
||||
export function createRelationExtractionFlow({
|
||||
maxBatchSize,
|
||||
maxConcurrency = 4,
|
||||
maxRelationsPerNode = 100,
|
||||
model,
|
||||
nodes,
|
||||
@ -72,6 +75,10 @@ export function createRelationExtractionFlow({
|
||||
throw new Error("Relation extraction maxBatchSize must be at least 1");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxConcurrency) || maxConcurrency < 1) {
|
||||
throw new Error("Relation extraction maxConcurrency must be at least 1");
|
||||
}
|
||||
|
||||
if (!Number.isInteger(maxRelationsPerNode) || maxRelationsPerNode < 1) {
|
||||
throw new Error("Relation extraction maxRelationsPerNode must be at least 1");
|
||||
}
|
||||
@ -113,34 +120,32 @@ export function createRelationExtractionFlow({
|
||||
};
|
||||
}
|
||||
|
||||
const generated = await Promise.all(
|
||||
orderedNodes.map(async (node) => {
|
||||
const entities = extractedEntitiesFromNodeMetadata(node);
|
||||
const result = await provider.extract({
|
||||
entities,
|
||||
maxRelations: maxRelationsPerNode,
|
||||
model,
|
||||
node: cloneKnowledgeNode(node),
|
||||
prompt: relationExtractionPrompt(node, entities),
|
||||
promptVersion,
|
||||
...(tenantId ? { tenantId } : {}),
|
||||
});
|
||||
const relations = validateExtractedRelations(result.relations, maxRelationsPerNode);
|
||||
const generated = await mapWithConcurrency(orderedNodes, maxConcurrency, async (node) => {
|
||||
const entities = extractedEntitiesFromNodeMetadata(node);
|
||||
const result = await provider.extract({
|
||||
entities,
|
||||
maxRelations: maxRelationsPerNode,
|
||||
model,
|
||||
node: cloneKnowledgeNode(node),
|
||||
prompt: relationExtractionPrompt(node, entities),
|
||||
promptVersion,
|
||||
...(tenantId ? { tenantId } : {}),
|
||||
});
|
||||
const relations = validateExtractedRelations(result.relations, maxRelationsPerNode);
|
||||
|
||||
return {
|
||||
id: node.id,
|
||||
metadata: relationExtractionMetadata({
|
||||
metadata: result.metadata,
|
||||
model,
|
||||
node,
|
||||
now,
|
||||
promptVersion,
|
||||
relations,
|
||||
traceId,
|
||||
}),
|
||||
};
|
||||
}),
|
||||
);
|
||||
return {
|
||||
id: node.id,
|
||||
metadata: relationExtractionMetadata({
|
||||
metadata: result.metadata,
|
||||
model,
|
||||
node,
|
||||
now,
|
||||
promptVersion,
|
||||
relations,
|
||||
traceId,
|
||||
}),
|
||||
};
|
||||
});
|
||||
const extractedNodes = await nodes.updateMetadataMany({
|
||||
knowledgeSpaceId,
|
||||
patches: generated,
|
||||
@ -193,13 +198,7 @@ function validateExtractedRelations(
|
||||
relations: readonly ExtractedRelation[],
|
||||
maxRelationsPerNode: number,
|
||||
): ExtractedRelation[] {
|
||||
if (relations.length > maxRelationsPerNode) {
|
||||
throw new Error(
|
||||
`Relation extraction provider returned ${relations.length} relations over maxRelationsPerNode=${maxRelationsPerNode}`,
|
||||
);
|
||||
}
|
||||
|
||||
return relations.map((relation) => {
|
||||
const validated = relations.map((relation, index) => {
|
||||
if (!RELATION_EXTRACTION_TYPES.has(relation.type)) {
|
||||
throw new Error("Relation extraction relation type is unsupported");
|
||||
}
|
||||
@ -221,13 +220,24 @@ function validateExtractedRelations(
|
||||
}
|
||||
|
||||
return {
|
||||
confidence: relation.confidence,
|
||||
...(relation.metadata ? { metadata: cloneJsonObject(relation.metadata) } : {}),
|
||||
object: relation.object.trim(),
|
||||
subject: relation.subject.trim(),
|
||||
type: relation.type,
|
||||
index,
|
||||
relation: {
|
||||
confidence: relation.confidence,
|
||||
...(relation.metadata ? { metadata: cloneJsonObject(relation.metadata) } : {}),
|
||||
object: relation.object.trim(),
|
||||
subject: relation.subject.trim(),
|
||||
type: relation.type,
|
||||
},
|
||||
};
|
||||
});
|
||||
|
||||
return validated
|
||||
.sort(
|
||||
(left, right) =>
|
||||
right.relation.confidence - left.relation.confidence || left.index - right.index,
|
||||
)
|
||||
.slice(0, maxRelationsPerNode)
|
||||
.map(({ relation }) => relation);
|
||||
}
|
||||
|
||||
function relationExtractionMetadata({
|
||||
|
||||
@ -0,0 +1,4 @@
|
||||
import { createConcurrencyGate } from "./bounded-concurrency";
|
||||
|
||||
/** Protects the shared model runtime across all entity and relation extraction providers. */
|
||||
export const semanticExtractionModelRequestGate = createConcurrencyGate(4);
|
||||
@ -12,6 +12,10 @@ import {
|
||||
import { createExtractionQualityControlFlow } from "./extraction-quality-control-flow";
|
||||
import { createInMemoryGraphIndexRepository } from "./graph-index-repository";
|
||||
import { createInMemoryKnowledgeNodeRepository } from "./knowledge-node-repository";
|
||||
import {
|
||||
type RelationExtractionProvider,
|
||||
createRelationExtractionFlow,
|
||||
} from "./relation-extraction-flow";
|
||||
import { createSemanticIngestionPostProcessor } from "./semantic-ingestion-postprocessor";
|
||||
|
||||
const knowledgeSpaceId = "018f0d60-7a49-7cc2-9c1b-5b36f18f2c42";
|
||||
@ -234,6 +238,86 @@ describe("createSemanticIngestionPostProcessor", () => {
|
||||
}),
|
||||
).rejects.toThrow("Semantic ingestion node count exceeds maxNodesPerArtifact=1");
|
||||
});
|
||||
|
||||
it("bounds concurrent entity and relation model requests", async () => {
|
||||
const nodes = createInMemoryKnowledgeNodeRepository({
|
||||
maxBatchSize: 10,
|
||||
maxListLimit: 10,
|
||||
maxNodes: 10,
|
||||
});
|
||||
const graph = createInMemoryGraphIndexRepository({
|
||||
maxBatchSize: 10,
|
||||
maxEntities: 20,
|
||||
maxRelations: 20,
|
||||
});
|
||||
await nodes.createMany(
|
||||
Array.from({ length: 6 }, (_, index) =>
|
||||
semanticNode(
|
||||
`018f0d60-7a49-7cc2-9c1b-5b36f18f2c${String(index + 1).padStart(2, "0")}`,
|
||||
index * 10,
|
||||
),
|
||||
),
|
||||
);
|
||||
let activeEntityCalls = 0;
|
||||
let maxActiveEntityCalls = 0;
|
||||
let activeRelationCalls = 0;
|
||||
let maxActiveRelationCalls = 0;
|
||||
const entityProvider: EntityExtractionProvider = {
|
||||
extract: async () => {
|
||||
activeEntityCalls += 1;
|
||||
maxActiveEntityCalls = Math.max(maxActiveEntityCalls, activeEntityCalls);
|
||||
await Promise.resolve();
|
||||
activeEntityCalls -= 1;
|
||||
|
||||
return {
|
||||
entities: [{ confidence: 0.9, text: "Acme", type: "organization" }],
|
||||
};
|
||||
},
|
||||
};
|
||||
const relationProvider: RelationExtractionProvider = {
|
||||
extract: async () => {
|
||||
activeRelationCalls += 1;
|
||||
maxActiveRelationCalls = Math.max(maxActiveRelationCalls, activeRelationCalls);
|
||||
await Promise.resolve();
|
||||
activeRelationCalls -= 1;
|
||||
|
||||
return { relations: [] };
|
||||
},
|
||||
};
|
||||
const processor = createSemanticIngestionPostProcessor({
|
||||
entityExtraction: createEntityExtractionFlow({
|
||||
maxBatchSize: 10,
|
||||
maxConcurrency: 2,
|
||||
model: "entity-llm",
|
||||
nodes,
|
||||
provider: entityProvider,
|
||||
}),
|
||||
extractionQuality: createExtractionQualityControlFlow({
|
||||
maxBatchSize: 10,
|
||||
nodes,
|
||||
}),
|
||||
graph,
|
||||
maxNodesPerArtifact: 10,
|
||||
nodes,
|
||||
relationExtraction: createRelationExtractionFlow({
|
||||
maxBatchSize: 10,
|
||||
maxConcurrency: 2,
|
||||
model: "relation-llm",
|
||||
nodes,
|
||||
provider: relationProvider,
|
||||
}),
|
||||
});
|
||||
|
||||
await expect(
|
||||
processor.process({
|
||||
knowledgeSpaceId,
|
||||
parseArtifact: { id: parseArtifactId },
|
||||
tenantId: "tenant-1",
|
||||
}),
|
||||
).resolves.toMatchObject({ nodesScanned: 6, nodesUpdated: 6 });
|
||||
expect(maxActiveEntityCalls).toBe(2);
|
||||
expect(maxActiveRelationCalls).toBe(2);
|
||||
});
|
||||
});
|
||||
|
||||
function createRecordingEntityProvider(): EntityExtractionProvider & {
|
||||
|
||||
@ -476,6 +476,57 @@ describe("createDifyModelRuntimeClient", () => {
|
||||
});
|
||||
});
|
||||
|
||||
it("converts LLM stream read timeouts into retryable runtime errors", async () => {
|
||||
const client = createDifyModelRuntimeClient({
|
||||
apiKey: "inner-secret",
|
||||
baseUrl: "http://api:5001",
|
||||
fetch: vi.fn(async (_input, init) => {
|
||||
const signal = init?.signal;
|
||||
if (!signal) {
|
||||
throw new Error("missing signal");
|
||||
}
|
||||
|
||||
return new Response(
|
||||
new ReadableStream<Uint8Array>({
|
||||
start(controller) {
|
||||
signal.addEventListener("abort", () => controller.error(signal.reason), {
|
||||
once: true,
|
||||
});
|
||||
},
|
||||
}),
|
||||
);
|
||||
}),
|
||||
requestTimeoutMs: 5,
|
||||
});
|
||||
|
||||
await expect(collect(client.invokeLlm(llmInput()))).rejects.toMatchObject({
|
||||
code: "dify_model_runtime_timeout",
|
||||
retryable: true,
|
||||
});
|
||||
});
|
||||
|
||||
it("converts LLM stream transport failures into retryable request errors", async () => {
|
||||
const client = createDifyModelRuntimeClient({
|
||||
apiKey: "inner-secret",
|
||||
baseUrl: "http://api:5001",
|
||||
fetch: vi.fn(
|
||||
async () =>
|
||||
new Response(
|
||||
new ReadableStream<Uint8Array>({
|
||||
start(controller) {
|
||||
controller.error(new TypeError("terminated"));
|
||||
},
|
||||
}),
|
||||
),
|
||||
),
|
||||
});
|
||||
|
||||
await expect(collect(client.invokeLlm(llmInput()))).rejects.toMatchObject({
|
||||
code: "dify_model_runtime_request_failed",
|
||||
retryable: true,
|
||||
});
|
||||
});
|
||||
|
||||
it("rejects missing, malformed, and oversized unary response bodies", async () => {
|
||||
const cases: ReadonlyArray<[Response, number | undefined, string]> = [
|
||||
[new Response(null), undefined, "dify_model_runtime_response_invalid"],
|
||||
|
||||
@ -232,6 +232,18 @@ export function createDifyModelRuntimeClient(
|
||||
}
|
||||
yield unwrapEnvelope(envelope.data);
|
||||
}
|
||||
} catch (cause) {
|
||||
if (deadline.signal.aborted) {
|
||||
throw deadlineError(deadline.signal.reason);
|
||||
}
|
||||
if (cause instanceof DifyModelRuntimeError) {
|
||||
throw cause;
|
||||
}
|
||||
throw new DifyModelRuntimeError("Dify model runtime stream failed", {
|
||||
cause,
|
||||
code: "dify_model_runtime_request_failed",
|
||||
retryable: true,
|
||||
});
|
||||
} finally {
|
||||
deadline.cleanup();
|
||||
}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user