From 1bf7879c1e89bd1d423c25a2975cefd8c16547ce Mon Sep 17 00:00:00 2001 From: Ajax Davis Date: Thu, 4 Dec 2025 13:59:22 +1000 Subject: [PATCH] feat: implement BM25 search with context from last 3 user messages MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Implement proper BM25 scoring algorithm in /api/tools/search - Term frequency with saturation (k1 = 1.5) - Length normalization (b = 0.75) - Inverse document frequency (IDF) - Accept recent messages via 'messages' query param for better context - Update search-registry tool to: - Use /api/tools/search endpoint (not /api/tools) - Pass last 3 user messages for contextual search - Include recentMessages in tool input schema - Update chat API to extract and pass last 3 user messages to search BM25 formula: Σ IDF(qi) * (tf * (k1 + 1)) / (tf + k1 * (1 - b + b * |D| / avgdl)) 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude --- apps/playground/src/app/api/chat/route.ts | 20 +++- apps/web/src/app/api/tools/search/route.ts | 103 ++++++++++++++++---- packages/tools/search-registry/src/index.ts | 54 +++++----- 3 files changed, 125 insertions(+), 52 deletions(-) diff --git a/apps/playground/src/app/api/chat/route.ts b/apps/playground/src/app/api/chat/route.ts index 5f390ce..4ef7247 100644 --- a/apps/playground/src/app/api/chat/route.ts +++ b/apps/playground/src/app/api/chat/route.ts @@ -74,7 +74,7 @@ export async function POST(request: NextRequest) { execute: typeof searchTpmjsToolsTool.execute, }); - // 2. Extract user query from last message for tool search + // 2. Extract user query and last 3 user messages for tool search const lastMessage = messages[messages.length - 1]; let userQuery = ''; if (lastMessage?.role === 'user') { @@ -88,7 +88,24 @@ export async function POST(request: NextRequest) { } } + // Get last 3 user messages for context + const recentUserMessages = messages + .filter((msg) => msg.role === 'user') + .slice(-3) + .map((msg) => { + // Extract text from parts + const parts = (msg as any).parts || []; + for (const part of parts) { + if (part.type === 'text') { + return part.text; + } + } + return ''; + }) + .filter(Boolean); + console.log(`💬 User query: "${userQuery}"`); + console.log(`📝 Recent messages: ${recentUserMessages.length}`); // 3. Automatically search for relevant tools based on the user's message if (userQuery && userQuery.trim().length > 0) { @@ -100,6 +117,7 @@ export async function POST(request: NextRequest) { { query: userQuery, limit: 5, // Get top 5 relevant tools + recentMessages: recentUserMessages, }, {} as any ); diff --git a/apps/web/src/app/api/tools/search/route.ts b/apps/web/src/app/api/tools/search/route.ts index 9e543ad..722b722 100644 --- a/apps/web/src/app/api/tools/search/route.ts +++ b/apps/web/src/app/api/tools/search/route.ts @@ -5,23 +5,51 @@ export const runtime = 'nodejs'; export const dynamic = 'force-dynamic'; export const maxDuration = 60; -// Simple text-based search scoring (fallback until BM25 is fixed) -function calculateTextScore(query: string, document: string): number { - const queryTokens = query.toLowerCase().split(/\s+/); - const docLower = document.toLowerCase(); +// BM25 parameters +const k1 = 1.5; // term frequency saturation parameter +const b = 0.75; // length normalization parameter + +// Tokenize text into words +function tokenize(text: string): string[] { + return text + .toLowerCase() + .replace(/[^\w\s]/g, ' ') + .split(/\s+/) + .filter((t) => t.length > 0); +} + +// Calculate term frequency +function termFrequency(term: string, tokens: string[]): number { + return tokens.filter((t) => t === term).length; +} + +// Calculate BM25 score +function calculateBM25( + query: string, + document: string, + avgDocLength: number, + totalDocs: number, + docFrequencies: Map +): number { + const queryTokens = tokenize(query); + const docTokens = tokenize(document); + const docLength = docTokens.length; let score = 0; - for (const token of queryTokens) { - // Exact match in document - if (docLower.includes(token)) { - score += 1; - } + for (const term of queryTokens) { + const tf = termFrequency(term, docTokens); + if (tf === 0) continue; - // Boost if token appears in beginning (likely more relevant) - if (docLower.startsWith(token)) { - score += 0.5; - } + // IDF calculation + const docFreq = docFrequencies.get(term) || 0; + const idf = Math.log((totalDocs - docFreq + 0.5) / (docFreq + 0.5) + 1); + + // BM25 formula + const numerator = tf * (k1 + 1); + const denominator = tf + k1 * (1 - b + b * (docLength / avgDocLength)); + + score += idf * (numerator / denominator); } return score; @@ -36,7 +64,13 @@ export async function GET(request: Request) { const category = searchParams.get('category'); const limit = Math.min(Number.parseInt(searchParams.get('limit') || '10'), 50); - console.log(`🔎 [SEARCH API] Query: "${query}", Category: ${category}, Limit: ${limit}`); + // Get recent messages for context (passed as JSON in 'messages' param) + const messagesParam = searchParams.get('messages'); + const recentMessages = messagesParam ? JSON.parse(messagesParam) : []; + + console.log( + `🔎 [SEARCH API] Query: "${query}", Category: ${category}, Limit: ${limit}, Messages: ${recentMessages.length}` + ); // Fetch all tools with package info const tools = await prisma.tool.findMany({ @@ -50,20 +84,47 @@ export async function GET(request: Request) { console.log(`📊 [SEARCH API] Found ${tools.length} tools in database`); - // Build searchable documents and calculate scores - const scoredResults = tools.map((tool) => { - const document = [ + // Combine query with recent messages for better context + const fullQuery = [query, ...recentMessages].filter(Boolean).join(' '); + console.log(`🔍 [SEARCH API] Full search context: "${fullQuery.slice(0, 100)}..."`); + + // Build all documents first + const documents = tools.map((tool) => ({ + tool, + text: [ tool.description, tool.exportName, tool.package.npmPackageName, tool.package.npmDescription || '', ...(tool.package.npmKeywords || []), - ].join(' '); + ].join(' '), + })); - const textScore = calculateTextScore(query, document); - const qualityBoost = Number.parseFloat(tool.qualityScore || '0') * 0.5; + // Calculate document frequencies (IDF) + const docFrequencies = new Map(); + const queryTokens = tokenize(fullQuery); + + for (const term of queryTokens) { + let count = 0; + for (const doc of documents) { + const docTokens = tokenize(doc.text); + if (docTokens.includes(term)) { + count++; + } + } + docFrequencies.set(term, count); + } + + // Calculate average document length + const totalTokens = documents.reduce((sum, doc) => sum + tokenize(doc.text).length, 0); + const avgDocLength = totalTokens / documents.length; + + // Calculate BM25 scores + const scoredResults = documents.map(({ tool, text }) => { + const bm25Score = calculateBM25(fullQuery, text, avgDocLength, tools.length, docFrequencies); + const qualityBoost = Number(tool.qualityScore ?? 0) * 0.5; const downloadBoost = Math.log10((tool.package.npmDownloadsLastMonth || 0) + 1) * 0.1; - const finalScore = textScore + qualityBoost + downloadBoost; + const finalScore = bm25Score + qualityBoost + downloadBoost; return { tool, score: finalScore }; }); diff --git a/packages/tools/search-registry/src/index.ts b/packages/tools/search-registry/src/index.ts index 8cdfc70..2e488c4 100644 --- a/packages/tools/search-registry/src/index.ts +++ b/packages/tools/search-registry/src/index.ts @@ -7,6 +7,7 @@ type SearchTpmjsToolsInput = { query: string; category?: string; limit?: number; + recentMessages?: string[]; }; /** @@ -51,17 +52,30 @@ export const searchTpmjsToolsTool = tool({ minimum: 1, maximum: 20, }, + recentMessages: { + type: 'array', + description: 'Recent user messages for context (optional)', + items: { + type: 'string', + }, + }, }, required: ['query'], additionalProperties: false, }), - async execute({ query, category, limit = 10 }) { - console.log('🔍 searchTpmjsTools.execute() called with:', { query, category, limit }); + async execute({ query, category, limit = 10, recentMessages = [] }) { + console.log('🔍 searchTpmjsTools.execute() called with:', { + query, + category, + limit, + recentMessages: recentMessages.length, + }); const params = new URLSearchParams({ q: query, limit: String(limit), ...(category && { category }), + ...(recentMessages.length > 0 && { messages: JSON.stringify(recentMessages) }), }); // Use environment variable, or default to production API (falls back to localhost in dev) @@ -69,9 +83,8 @@ export const searchTpmjsToolsTool = tool({ process.env.TPMJS_API_URL || (process.env.NODE_ENV === 'production' ? 'https://tpmjs.com' : 'http://localhost:3000'); - // Note: /api/tools/search endpoint exists in local dev but not deployed yet - // Using /api/tools as fallback with client-side filtering for now - const url = `${baseUrl}/api/tools?${params}`; + // Use the /api/tools/search endpoint for BM25 scoring with context + const url = `${baseUrl}/api/tools/search?${params}`; console.log(`🌐 Fetching: ${url}`); @@ -83,37 +96,18 @@ export const searchTpmjsToolsTool = tool({ throw new Error(`Search failed: ${response.statusText}`); } + // biome-ignore lint/suspicious/noExplicitAny: API response types vary const data = (await response.json()) as any; console.log('📦 Search response data:', JSON.stringify(data, null, 2)); - // Handle both /api/tools (deployed) and /api/tools/search (local dev) responses - const toolsArray = data.results?.tools || data.data || []; - - // Client-side filtering if query is provided (since deployed /api/tools doesn't support search yet) - let filteredTools = toolsArray; - if (query?.trim()) { - const queryLower = query.toLowerCase(); - filteredTools = toolsArray.filter((tool: any) => { - const searchableText = [ - tool.description, - tool.exportName, - tool.package?.npmPackageName, - tool.package?.npmDescription, - ...(tool.package?.npmKeywords || []), - ] - .join(' ') - .toLowerCase(); - return searchableText.includes(queryLower); - }); - } - - // Apply limit - const limitedTools = filteredTools.slice(0, limit); + // Handle /api/tools/search response structure + const toolsArray = data.results?.tools || []; return { query, - matchCount: filteredTools.length, - tools: limitedTools.map((tool: any) => ({ + matchCount: toolsArray.length, + // biome-ignore lint/suspicious/noExplicitAny: Tool types from API vary + tools: toolsArray.map((tool: any) => ({ toolId: tool.id, packageName: tool.package.npmPackageName, exportName: tool.exportName,