跳转到正文

retrieval.ts ​

源文件: docs/worker-api/examples/retrieval.ts · 下载原文件 · 示例使用说明

ts
import { createDitto, createMemoryWorker, loadRuntimeConfigFile } from "@codesoul-co/ditto";
import type { MemorySearchProvider, MemoryStore, MemoryItem } from "@codesoul-co/ditto/worker/memory";
import {
  createRetrieval, createRetrievalWorker, RetrievalTargetRegistry, RetrievalError, retrievalSearchNode,
  embedContents, validateVector, createVectorSearchProvider, createTextSearchProvider,
  createHybridSearchProvider, createCosineReranker, rerankCandidates, createRerankSearchProvider,
  createHttpEmbeddingProvider, embeddingConfigFromEnv, createSqlSearchProvider, createMilvusSearchProvider,
  type RetrievalSearchProvider, type RetrievalSearchInput, type RetrievalSearchOutput,
  type EmbeddingProvider, type RerankProvider, type SqlSearchOptions, type MilvusSearchOptions,
} from "@codesoul-co/ditto-retrieval";
import {
  createMemoryRetrievalProvider, createRetrievalMemorySearchProvider,
  RemoteRetrievalSearchProvider, mapMemoryCandidates,
} from "@codesoul-co/ditto-retrieval/adapters/memory";

export const request: RetrievalSearchInput = { query: { content: "agent memory" }, target: { name: "kb" }, limit: 5 };

// example: setup
export function setupRetrieval(provider: RetrievalSearchProvider) {
  const config = loadRuntimeConfigFile("ditto.yaml", process.env);
  const providers = new RetrievalTargetRegistry({ kb: { defaultStrategy: "vector", providers: { vector: provider } } });
  const runtime = createDitto({ config, workers: [createRetrievalWorker({ providers, concurrency: 4 })] });
  const retrieval = createRetrieval({ providers, defaults: config.retrieval });
  return { runtime, retrieval };
}

// example: search
export async function retrievalSearch(provider: RetrievalSearchProvider) {
  const { runtime, retrieval } = setupRetrieval(provider);
  try {
    const local = await retrieval.search(request);
    const routed = await runtime.invoke("RETRIEVAL.SEARCH", request);
    if (routed.status !== "success" || !routed.output) throw new Error(routed.error?.code ?? routed.status);
    return { local, candidates: routed.output.candidates, target: routed.output.target };
  } finally { await runtime.close(); }
}

// example: registry
export async function registryResolve(provider: RetrievalSearchProvider) {
  const providers = new RetrievalTargetRegistry({ kb: { defaultStrategy: "keyword", providers: { keyword: provider } } });
  const selected = providers.resolve({ name: "kb" });
  return selected.search({ ...request, limit: 5 }); // Raw output, strategy is filled with "keyword".
}

// example: cancel
export async function cancelRetrieval(provider: RetrievalSearchProvider) {
  const providers = new RetrievalTargetRegistry({ kb: { defaultStrategy: "custom", providers: { custom: provider } } });
  const retrieval = createRetrieval({ providers, defaults: { searchLimit: 5 } });
  const controller = new AbortController(); controller.abort();
  return retrieval.search(request, { searchLimit: 10 }, { signal: controller.signal });
}

// example: embedding
export async function embeddingApis(embedding: EmbeddingProvider, dimensions: number) {
  const direct = await embedding.embed({ contents: ["query text"], purpose: "query" });
  const documents = await embedContents(embedding, { contents: ["document A", "document B"], purpose: "document" }, { batchSize: 32, dimensions });
  validateVector(documents[0], dimensions); // Returns void; throws on invalid/empty/mismatched vectors.
  return { direct, documents };
}

// example: httpEmbedding
export async function httpEmbedding() {
  const config = loadRuntimeConfigFile("ditto.yaml", process.env);
  const runtime = createDitto({ config });
  try {
    const embedding = createHttpEmbeddingProvider({ ...embeddingConfigFromEnv(process.env), sandbox: runtime.services.sandbox, timeoutMs: config.timeoutMs });
    return await embedContents(embedding, { contents: ["query text"], purpose: "query" }, {}, { defaults: config.retrieval });
  } finally { await runtime.close(); }
}

// example: vector
export async function externalVector(backend: RetrievalSearchProvider, embedding: EmbeddingProvider) {
  const vector = createVectorSearchProvider({ backend, embedding });
  return vector.search(request);
}
export async function nativeVector(backend: RetrievalSearchProvider) {
  return createVectorSearchProvider({ backend, nativeEmbedding: true }).search(request);
}
export async function precomputedVector(backend: RetrievalSearchProvider, vector: readonly number[]) {
  return createVectorSearchProvider({ backend, dimensions: vector.length }).search({ ...request, query: { content: vector } });
}

// example: text
export async function textSearch(backend: RetrievalSearchProvider) {
  return createTextSearchProvider(backend).search({ ...request, strategy: "keyword" });
}

// example: hybrid
export async function hybridSearch(vector: RetrievalSearchProvider, text: RetrievalSearchProvider) {
  const hybrid = createHybridSearchProvider({ candidateLimit: 50, rrfK: 60,
    branches: [{ provider: vector, strategy: "vector", weight: 1 }, { provider: text, strategy: "keyword", weight: 0.7 }],
  });
  return hybrid.search({ ...request, strategy: "hybrid" });
}

// example: rerank
export async function rerankApis(embedding: EmbeddingProvider, output: RetrievalSearchOutput) {
  const reranker = createCosineReranker(embedding);
  const input = { query: request.query, candidates: output.candidates, limit: Math.max(1, Math.min(3, output.candidates.length)) };
  const order = await reranker.rank(input); // Raw zero-based indexes and scores.
  const candidates = await rerankCandidates(reranker, input); // Original candidates in ranked order.
  return { order, candidates };
}
export async function rerankSearch(search: RetrievalSearchProvider, reranker: RerankProvider) {
  return createRerankSearchProvider({ search, reranker, candidateLimit: 50 }).search(request);
}

// example: sql
interface SqlRow { id: string; content: string; score: number; }
export function sqlSearch(query: SqlSearchOptions<SqlRow>["query"]) {
  return createSqlSearchProvider<SqlRow>({
    prepare(input) {
      if (typeof input.query.content !== "string" || !input.target.namespace || input.filter || input.options) {
        throw new RetrievalError("RETRIEVAL_INVALID_INPUT", "Expected text and namespace without additional filters or options");
      }
      return {
        text: "SELECT id, content, ts_rank_cd(to_tsvector('english', content), plainto_tsquery('english', $1)) AS score FROM memories WHERE tenant = $2 AND to_tsvector('english', content) @@ plainto_tsquery('english', $1) ORDER BY score DESC, id ASC LIMIT $3",
        values: [input.query.content, input.target.namespace, input.limit],
      };
    },
    query, // e.g. async ({ text, values }) => (await pgPool.query(text, [...values])).rows
    mapRow: row => ({ id: row.id, content: row.content, score: row.score }),
  });
}

// example: milvus
interface MilvusHit { id: string | number; score: number; memory: MemoryItem; }
export function milvusSearch(search: MilvusSearchOptions<MilvusHit>["search"]) {
  return createMilvusSearchProvider<MilvusHit>({ collection: "memories", vectorField: "embedding", outputFields: ["memory"],
    search, // Adapt the application's SDK response to { status, results }.
    scope(input) {
      if (!input.target.namespace || input.filter || input.options) throw new RetrievalError("RETRIEVAL_INVALID_INPUT", "Expected a namespace without extra filters or options");
      return { filter: "tenant == {tenant}", exprValues: { tenant: input.target.namespace } };
    },
    mapHit: hit => ({ id: String(hit.id), content: hit.memory, score: hit.score }),
  });
}

// example: memoryNative
export async function nativeMemoryInRetrieval(nativeSearch: MemorySearchProvider) {
  const provider = createMemoryRetrievalProvider(nativeSearch);
  const output = await provider.search(request);
  return mapMemoryCandidates(output); // content must be a complete MemoryItem, not just text.
}

// example: memoryNamespace
export function nativeMemoryWithNamespace(nativeSearch: MemorySearchProvider) {
  return createMemoryRetrievalProvider(nativeSearch, input => {
    if (!input.target.namespace) throw new RetrievalError("RETRIEVAL_INVALID_INPUT", "A namespace is required");
    return {
      query: input.query.content, filter: { ...input.filter, tenant: input.target.namespace },
      ...(input.limit === undefined ? {} : { limit: input.limit }),
      ...(input.strategy === undefined ? {} : { strategy: input.strategy }),
      ...(input.options === undefined ? {} : { options: input.options }),
    };
  }); // The native plugin must implement this tenant filter and enforce caller authorization.
}

// example: memoryLocal
export async function localMemoryPipeline(store: MemoryStore, provider: RetrievalSearchProvider) {
  const search = createRetrievalMemorySearchProvider({ provider, target: { name: "kb" }, strategy: "vector", defaults: { searchLimit: 5 } });
  const runtime = createDitto({ workers: [createMemoryWorker({ store, search })] });
  try { return await runtime.invoke("MEMORY.SEARCH", { query: "agent memory", limit: 3 }); }
  finally { await runtime.close(); }
}

// example: memoryRemote
export async function delegatedMemory(store: MemoryStore, provider: RetrievalSearchProvider) {
  const { runtime } = setupRetrieval(provider);
  const search = new RemoteRetrievalSearchProvider({ runtime, target: { name: "kb" } });
  runtime.register(createMemoryWorker({ store, search }));
  try {
    const raw = await search.search({ query: "agent memory", limit: 3 }); // Raw MemorySearchOutput.
    const routed = await runtime.invoke("MEMORY.SEARCH", { query: "agent memory", limit: 3 });
    return { raw, routed };
  } finally { await runtime.close(); }
}

// example: memoryMap
export function mapTextCandidates(output: RetrievalSearchOutput) {
  return output.candidates.map(candidate => {
    if (!candidate.id) throw new Error("The application requires candidate ids");
    return { memory: { id: candidate.id, content: candidate.content }, ...(candidate.score === undefined ? {} : { score: candidate.score }) };
  }); // Use as mapOutput only when content is the application's complete memory content.
}

// example: scaffold
export function retrievalDescriptor(provider: RetrievalSearchProvider) {
  const providers = new RetrievalTargetRegistry({ kb: { defaultStrategy: "custom", providers: { custom: provider } } });
  const retrieval = createRetrieval({ providers });
  return retrievalSearchNode.define("RETRIEVAL", input => retrieval.search(input));
}

Ditto · @codesoul-co/ditto · Node.js 24+