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
16 changes: 16 additions & 0 deletions packages/core/src/conversations/threads/message-text.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import type { CollaborationPart, Message, ToolCallPart } from "@sugabots/contracts";
import { Option, Schema } from "effect";
import { TOOL_SEARCH } from "../tools/tool-search/tool.ts";

/** The tool an agent searches its thread's older history with. */
export const SEARCH_HISTORY_TOOL = "search_history";
Expand Down Expand Up @@ -43,6 +45,7 @@ export function describeToolCall(call: ToolCallPart): string {
case "failed":
return `${asked}; it failed: ${call.error ?? "no reason given"}]`;
case "completed":
if (call.tool === TOOL_SEARCH) return `${asked}: found ${foundToolNames(call.output)}]`;
return `${asked}: ${clipped(
JSON.stringify(call.output),
call.tool === SEARCH_HISTORY_TOOL
Expand All @@ -52,6 +55,19 @@ export function describeToolCall(call: ToolCallPart): string {
}
}

const FoundToolNames = Schema.Struct({
tools: Schema.Array(Schema.Struct({ tool: Schema.String })),
});

/** foundToolNames leaves out schemas: a later turn that needs one searches again. */
function foundToolNames(output: unknown): string {
const found = Option.match(Schema.decodeUnknownOption(FoundToolNames)(output), {
onNone: () => [],
onSome: ({ tools }) => tools.map(({ tool }) => tool),
});
return found.length > 0 ? found.join(", ") : "no tools";
}

function clipped(text: string, limit: number): string {
return text.length <= limit ? text : `${text.slice(0, limit)}… (${text.length} characters)`;
}
Expand Down
2 changes: 1 addition & 1 deletion packages/core/src/conversations/threads/threads.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -986,7 +986,7 @@ describe.skipIf(!process.env.DATABASE_URL)("threads, against Postgres", async ()
const prompt = modelPrompt(turn.context, {
now: new Date(),
builtInTools: [],
connectionTools: [],
connectionTools: undefined,
});
expect(prompt.messages[0]?.content).toContain("The family is planning a trip.");
await turnRecords.complete(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -396,7 +396,7 @@ describe.skipIf(!process.env.DATABASE_URL)("collaboration, against Postgres", as
const promptMessages = modelPrompt(next.context, {
now: new Date(),
builtInTools: [],
connectionTools: [],
connectionTools: undefined,
}).messages;
expect(promptMessages).toContainEqual(
expect.objectContaining({
Expand Down
18 changes: 12 additions & 6 deletions packages/core/src/conversations/tools/connections.ts
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ import {
connectionToolKey,
connectionToolMutating,
} from "@sugabots/contracts";
import type { Tool } from "ai";
import type { JSONSchema7, Tool } from "ai";
import { Context, Effect, Layer } from "effect";
import type { Database } from "../../database/database.ts";
import { ConnectionRepository } from "../../providers/connections/connection-repository.ts";
Expand All @@ -27,10 +27,9 @@ import {
* asked for its tools, and closed when the turn ends. Each tool is keyed by
* the connection's handle and its own name, `linear__list_issues`, and
* carries whether it may change something, which decides how a failed turn
* after it is treated, and what the pod's bots may do with it. A tool set to
* `off` is still offered, so turning one off or on leaves the tools the model
* is sent, and the provider's cache of them, as they were; its calls are
* refused. A tool nobody has chosen for is treated as `toolAccessOf` says.
* after it is treated, and what the pod's bots may do with it. A call to a
* tool set to `off` is refused. A tool nobody has chosen for is treated as
* `toolAccessOf` says.
* A connection whose every tool is off is not opened at all.
*
* A server that cannot be reached is left out of the turn, with a line in the
Expand All @@ -40,6 +39,10 @@ import {

export interface OfferedTool {
tool: Tool;
handle: string;
description: string;
/** The server's schema, before the client adapts it into `tool`. */
inputSchema: JSONSchema7;
/** Whether a call may change something at the other end. */
mutating: boolean;
/** Whether each call runs, waits for a person to allow it, or is refused. */
Expand Down Expand Up @@ -122,9 +125,12 @@ export function from({
);
try {
const tools: Record<string, OfferedTool> = {};
for (const { described, tool } of await session.tools()) {
for (const { described, inputSchema, tool } of await session.tools()) {
tools[connectionToolKey(target.handle, described.name)] = {
tool,
handle: target.handle,
description: described.description ?? "",
inputSchema,
mutating: connectionToolMutating(described),
access: toolAccessOf(target.toolAccess, described),
connectionId: target.connectionId,
Expand Down
68 changes: 68 additions & 0 deletions packages/core/src/conversations/tools/tool-search/catalog.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
import { describe, expect, it } from "vitest";
import { type CatalogEntry, catalogListing, LISTING_CHARACTERS, searchCatalog } from "./catalog.ts";

const entry = (
handle: string,
name: string,
{ description = "", properties = {} as Record<string, object>, required = [] as string[] } = {},
): CatalogEntry => ({
key: `${handle}__${name}`,
handle,
remoteToolName: name,
description,
inputSchema: { type: "object", properties, required },
});

describe("the listing of connection tools", () => {
it("stays within its budget however many tools its connections have, and says how many it left out", () => {
const many = (handle: string) =>
Array.from({ length: 371 }, (_, index) =>
entry(handle, `run_report_${String(index).padStart(3, "0")}`),
);
const listing = catalogListing([...many("reports"), ...many("sales")]);

expect(listing.length).toBeLessThanOrEqual(LISTING_CHARACTERS);
expect(listing).toMatch(
/^- reports \(371 tools\): reports__run_report_000, .*, and \d+ more$/m,
);
expect(listing).toMatch(/^- sales \(371 tools\): sales__run_report_000, .*, and \d+ more$/m);
});

it("names every tool of a small connection by the full name it is called by", () => {
expect(catalogListing([entry("notes", "list_notes")])).toBe(
"- notes (1 tool): notes__list_notes",
);
});
});

describe("searching connection tools", () => {
it("finds a tool by a plural of a word in its name, ignoring words every description has", () => {
const entries = [
entry("tracker", "list_issue", { description: "Lists the issues in a project." }),
entry("tracker", "create_project", { description: "Creates a project for the team." }),
];

expect(searchCatalog(entries, "the open issues").map((found) => found.tool)).toEqual([
"tracker__list_issue",
]);
});

it("gives the best matches' schemas, and only what the others require", () => {
const entries = ["run_report", "run_report_export", "run_report_schedule"].map((name) =>
entry("reports", name, {
description: "Runs a report. Takes a while.",
properties: { query: { type: "string" } },
required: ["query"],
}),
);

const found = searchCatalog(entries, "run report");

expect(found.slice(0, 2).every((match) => "inputSchema" in match)).toBe(true);
expect(found[2]).toEqual({
tool: "reports__run_report_schedule",
description: "Runs a report.",
required: ["query"],
});
});
});
138 changes: 138 additions & 0 deletions packages/core/src/conversations/tools/tool-search/catalog.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
import type { JSONSchema7 } from "ai";
import type { OfferedTool } from "../connections.ts";

const SEARCH_RESULT_LIMIT = 5;

/** Only the best matches carry whole schemas, so a search adds about 3,000 tokens at most. */
const SCHEMAS_PER_SEARCH = 2;
const SEARCH_SCHEMA_CHARACTERS = 12_000;

/** About 2,000 tokens of the turn's note. */
export const LISTING_CHARACTERS = 8_000;

export type CatalogEntry = Pick<
OfferedTool,
"handle" | "remoteToolName" | "description" | "inputSchema"
> & { key: string };

/** buildCatalog returns the tools that may be called, sorted so the listing is stable across turns. */
export function buildCatalog(tools: Readonly<Record<string, OfferedTool>>): CatalogEntry[] {
return Object.entries(tools)
.filter(([, offered]) => offered.access !== "off")
.map(([key, offered]) => ({ ...offered, key }))
.sort((a, b) => (a.key < b.key ? -1 : a.key > b.key ? 1 : 0));
}

type FoundTool =
| { tool: string; description: string; inputSchema: JSONSchema7 }
| { tool: string; description: string; required: string[] };

/**
* searchCatalog returns the entries that match most of `query`'s words, a
* match in the tool's name weighing most. Ties go by key, so it is repeatable.
*/
export function searchCatalog(catalog: readonly CatalogEntry[], query: string): FoundTool[] {
const terms = [...new Set(splitWords(query))].filter((word) => !STOP_WORDS.has(word)).map(stem);
const ranked = catalog
.map((entry) => ({ entry, ...scoreMatch(entry, terms) }))
.filter(({ matched }) => matched > 0)
.sort(
(a, b) =>
b.matched - a.matched || b.weight - a.weight || (a.entry.key < b.entry.key ? -1 : 1),
)
.slice(0, SEARCH_RESULT_LIMIT);
let schemaCharacters = 0;
return ranked.map(({ entry }, rank) => {
schemaCharacters += JSON.stringify(entry.inputSchema).length;
if (rank < SCHEMAS_PER_SEARCH && schemaCharacters <= SEARCH_SCHEMA_CHARACTERS) {
return { tool: entry.key, description: entry.description, inputSchema: entry.inputSchema };
}
return {
tool: entry.key,
description: firstSentence(entry.description),
required: entry.inputSchema.required ?? [],
};
});
}

const STOP_WORDS = new Set(
"a an and are as at be by can do for from get i in is it me my of on or our that the this to we what when with you your".split(
" ",
),
);

function scoreMatch(
entry: CatalogEntry,
terms: readonly string[],
): { matched: number; weight: number } {
const name = stemWords(entry.remoteToolName);
const handle = stemWords(entry.handle);
const parameters = stemWords(Object.keys(entry.inputSchema.properties ?? {}).join(" "));
const description = stemWords(entry.description);
let matched = 0;
let weight = 0;
for (const term of terms) {
const termWeight =
(name.has(term) ? 3 : 0) +
(handle.has(term) ? 2 : 0) +
(parameters.has(term) ? 1 : 0) +
(description.has(term) ? 1 : 0);
if (termWeight > 0) matched += 1;
weight += termWeight;
}
return { matched, weight };
}

function stemWords(text: string): Set<string> {
return new Set(splitWords(text).map(stem));
}

function splitWords(text: string): string[] {
return text
.replace(/([a-z0-9])([A-Z])/g, "$1 $2")
.toLowerCase()
.split(/[^a-z0-9]+/)
.filter((word) => word.length > 0);
}

/** stem drops a plural ending, so `issues` finds `list_issue`. */
function stem(word: string): string {
if (word.length > 4 && word.endsWith("ies")) return `${word.slice(0, -3)}y`;
if (word.length > 3 && word.endsWith("s") && !word.endsWith("ss")) return word.slice(0, -1);
return word;
}

function firstSentence(text: string): string {
const end = text.search(/[.!?](\s|$)/);
return end < 0 ? text : text.slice(0, end + 1);
}

/**
* catalogListing names each connection's tools within an equal share of
* `LISTING_CHARACTERS`, so one with hundreds of tools cannot crowd out the rest.
*/
export function catalogListing(catalog: readonly CatalogEntry[]): string {
const byConnection = new Map<string, string[]>();
for (const entry of catalog) {
byConnection.set(entry.handle, [...(byConnection.get(entry.handle) ?? []), entry.key]);
}
const share = Math.floor(LISTING_CHARACTERS / Math.max(1, byConnection.size));
return [...byConnection]
.map(([connection, keys]) => {
const head = `- ${connection} (${keys.length} ${keys.length === 1 ? "tool" : "tools"})`;
const all = `${head}: ${keys.join(", ")}`;
// One character of each share is the line's break.
if (all.length < share) return all;
const more = `, and ${keys.length} more`;
const shown: string[] = [];
let length = head.length + 2 + more.length;
for (const key of keys) {
length += key.length + 2;
if (length >= share) break;
shown.push(key);
}
if (shown.length === 0) return head;
return `${head}: ${shown.join(", ")}, and ${keys.length - shown.length} more`;
})
.join("\n");
}
Loading
Loading