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
14 changes: 4 additions & 10 deletions runtime/ai/analyst_agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -183,16 +183,10 @@ func (t *AnalystAgent) Handler(ctx context.Context, args *AnalystAgentArgs) (*An
}

// Pre-invoke the load_skill tool for each analyst skill the user referenced in the prompt, on every invocation.
// A skill already loaded in this conversation is skipped: the model has it, and loading it again would repeat its whole body in the context.
loaded := loadedSkills(s)
for _, sk := range referencedSkills(args.Prompt, skillsForAgent(skills, parser.SkillAgentAnalyst)) {
if loaded[sk.Name] {
continue
}
_, err := s.CallTool(ctx, RoleAssistant, LoadSkillName, nil, &LoadSkillArgs{Name: sk.Name})
if err != nil && errors.Is(err, ctx.Err()) { // Don't exit on non-context errors
return nil, err
}
// The model sees the whole conversation, so a skill loaded in any earlier turn is skipped.
err = loadReferencedSkills(ctx, s, args.Prompt, skills, parser.SkillAgentAnalyst, loadedSkills(s))
if err != nil {
return nil, err
}

// Determine tools that can be used
Expand Down
6 changes: 6 additions & 0 deletions runtime/ai/developer_agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,12 @@ func (t *DeveloperAgent) Handler(ctx context.Context, args *DeveloperAgentArgs)
return nil, err
}
}
// Pre-invoke the load_skill tool for each developer skill the user referenced in the prompt.
// The model only sees this invocation's tool calls, so only the skills loaded above in this invocation are skipped.
err = loadReferencedSkills(ctx, s, args.Prompt, skills, parser.SkillAgentDeveloper, loadedSkills(s, FilterByParent(s.ParentID)))
if err != nil {
return nil, err
}

// Build initial completion messages
messages := []*aiv1.CompletionMessage{NewTextCompletionMessage(RoleSystem, systemPrompt)}
Expand Down
24 changes: 22 additions & 2 deletions runtime/ai/skill_references.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
package ai

import (
"context"
"errors"
"regexp"
"slices"
)
Expand All @@ -12,6 +14,22 @@ var (
chatReferenceAttrRegexp = regexp.MustCompile(`(\w+)="([^"]*)"`)
)

// loadReferencedSkills pre-invokes the load_skill tool for each of the agent's skills that the user referenced in the prompt, in order of first reference.
// A skill in loaded is skipped: the model has it, and loading it again would repeat its whole body in the context.
// Tool errors are recorded in the session and don't fail the agent; only context cancellation is returned.
func loadReferencedSkills(ctx context.Context, s *Session, prompt string, skills []*Skill, agent string, loaded map[string]bool) error {
for _, sk := range referencedSkills(prompt, skillsForAgent(skills, agent)) {
if loaded[sk.Name] {
continue
}
_, err := s.CallTool(ctx, RoleAssistant, LoadSkillName, nil, &LoadSkillArgs{Name: sk.Name})
if err != nil && errors.Is(err, ctx.Err()) {
return err
}
}
return nil
}

// referencedSkills returns the skills referenced in a prompt with a chat reference of type "skill", once each and in order of first reference.
// References to skills that are not in the given list are ignored.
func referencedSkills(prompt string, skills []*Skill) []*Skill {
Expand All @@ -35,10 +53,12 @@ func referencedSkills(prompt string, skills []*Skill) []*Skill {
}

// loadedSkills returns the names of the skills already loaded in the session, whether pre-invoked or called by the model.
// The predicates narrow the load_skill calls considered, e.g. to the calls of the current invocation when the model doesn't see earlier ones.
// A call whose result is an error, such as a skill that was not found yet, doesn't count: the model never got the body.
func loadedSkills(s *Session) map[string]bool {
func loadedSkills(s *Session, predicates ...Predicate) map[string]bool {
res := map[string]bool{}
for _, call := range s.Messages(FilterByType(MessageTypeCall), FilterByTool(LoadSkillName)) {
predicates = append(predicates, FilterByType(MessageTypeCall), FilterByTool(LoadSkillName))
for _, call := range s.Messages(predicates...) {
result, ok := s.Message(FilterByParent(call.ID), FilterByType(MessageTypeResult))
if !ok || result.ContentType == MessageContentTypeError {
continue
Expand Down
54 changes: 53 additions & 1 deletion runtime/ai/skill_references_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ func textTurn(text string) turnFunc {
}
}

// newSkillReferencesSession creates a session on a project with analyst, developer and always-apply skills, backed by the given simulated model.
// newSkillReferencesSession creates a session on a project with analyst and developer skills, one always-apply skill for each, backed by the given simulated model.
func newSkillReferencesSession(t *testing.T, script *scriptedAIService) *ai.Session {
s, _, _ := newSkillReferencesSessionWithRuntime(t, script)
return s
Expand Down Expand Up @@ -94,6 +94,7 @@ ARPU excludes trial users.`,
// Without agents, a skill only applies to the developer agent.
"skills/dev-conventions/SKILL.md": `---
description: Development conventions.
always_apply: true
---
Name models in snake_case.`,
},
Expand Down Expand Up @@ -245,3 +246,54 @@ func TestAnalystLoadsReferencedSkillsOnEveryTurn(t *testing.T) {
require.Len(t, script.inputs, 2)
require.Contains(t, loadSkillResults(script.inputs[1])["monthly-close"], "Compare revenue month over month.")
}

// TestDeveloperLoadsReferencedSkills verifies that the developer loads the skills referenced with a chat-reference tag in the prompt before the model's first turn,
// once per distinct developer skill, and ignores references to skills that don't exist or don't apply to the developer.
func TestDeveloperLoadsReferencedSkills(t *testing.T) {
script := &scriptedAIService{turns: []turnFunc{textTurn("done")}}
s := newSkillReferencesSession(t, script)

prompt := `<chat-reference>type="skill" skill="churn-review"</chat-reference> as a model, then ` +
`<chat-reference>skill="churn-review" type="skill"</chat-reference> again. ` +
`<chat-reference>type="skill" skill="does-not-exist"</chat-reference> ` +
`<chat-reference>type="skill" skill="monthly-close"</chat-reference> ` +
`<chat-reference>type="skill" skill="dev-conventions"</chat-reference>`
res, err := s.CallTool(t.Context(), ai.RoleUser, ai.DeveloperAgentName, nil, &ai.DeveloperAgentArgs{Prompt: prompt})
require.NoError(t, err)

// The always-apply dev-conventions is pre-loaded once and not again for its reference; the referenced developer skill is loaded once.
require.Equal(t, []string{"dev-conventions", "churn-review"}, loadedSkillNames(s, res.Call.ID))

// The referenced skill's body was in the model's input on its first turn.
require.NotEmpty(t, script.inputs)
results := loadSkillResults(script.inputs[0])
require.Contains(t, results["churn-review"], "List the countries with the most churned customers.")
require.NotContains(t, results, "monthly-close")
require.NotContains(t, results, "does-not-exist")
}

// TestDeveloperLoadsReferencedSkillsOnEveryTurn verifies that a skill referenced in a later turn is loaded in that turn,
// even if an earlier turn loaded it: the developer's model only sees the tool calls of the current turn.
func TestDeveloperLoadsReferencedSkillsOnEveryTurn(t *testing.T) {
script := &scriptedAIService{turns: []turnFunc{textTurn("first"), textTurn("second"), textTurn("third")}}
s := newSkillReferencesSession(t, script)

res1, err := s.CallTool(t.Context(), ai.RoleUser, ai.DeveloperAgentName, nil, &ai.DeveloperAgentArgs{Prompt: "Hello"})
require.NoError(t, err)
require.Equal(t, []string{"dev-conventions"}, loadedSkillNames(s, res1.Call.ID))

prompt := `<chat-reference>type="skill" skill="churn-review"</chat-reference> as a model`
res2, err := s.CallTool(t.Context(), ai.RoleUser, ai.DeveloperAgentName, nil, &ai.DeveloperAgentArgs{Prompt: prompt})
require.NoError(t, err)
require.Equal(t, []string{"dev-conventions", "churn-review"}, loadedSkillNames(s, res2.Call.ID))

res3, err := s.CallTool(t.Context(), ai.RoleUser, ai.DeveloperAgentName, nil, &ai.DeveloperAgentArgs{Prompt: prompt + " with a metrics view"})
require.NoError(t, err)
require.Equal(t, []string{"dev-conventions", "churn-review"}, loadedSkillNames(s, res3.Call.ID))

// The body was in the model's input on each of the later turns' first completion.
require.Len(t, script.inputs, 3)
require.NotContains(t, loadSkillResults(script.inputs[0]), "churn-review")
require.Contains(t, loadSkillResults(script.inputs[1])["churn-review"], "List the countries with the most churned customers.")
require.Contains(t, loadSkillResults(script.inputs[2])["churn-review"], "List the countries with the most churned customers.")
}
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
import { describe, expect, it } from "vitest";
import type { V1Resource } from "@rilldata/web-common/runtime-client";
import { ResourceKind } from "@rilldata/web-common/features/entity-management/resource-selectors.ts";
import { getSkillPickerItems } from "@rilldata/web-common/features/chat/core/context/picker/data/skills.ts";
import {
getSkillAgent,
getSkillPickerItems,
} from "@rilldata/web-common/features/chat/core/context/picker/data/skills.ts";
import { ToolName } from "@rilldata/web-common/features/chat/core/types.ts";

function skillResource(
name: string,
Expand All @@ -14,12 +18,26 @@ function skillResource(
};
}

describe("getSkillAgent", () => {
it("maps each chat agent to the agent named in skills", () => {
expect(getSkillAgent(ToolName.ANALYST_AGENT)).toBe("analyst");
expect(getSkillAgent(ToolName.DEVELOPER_AGENT)).toBe("developer");
});

it("has no skills for other agents", () => {
expect(getSkillAgent(ToolName.ROUTER_AGENT)).toBeUndefined();
});
});

describe("getSkillPickerItems", () => {
it("lists the skills for the analyst with their descriptions", () => {
const items = getSkillPickerItems([
skillResource("monthly-close", ["analyst"]),
skillResource("glossary", ["analyst", "developer"]),
]);
const items = getSkillPickerItems(
[
skillResource("monthly-close", ["analyst"]),
skillResource("glossary", ["analyst", "developer"]),
],
"analyst",
);
expect(
items.map((i) => [i.context.skill, i.context.label, i.description]),
).toEqual([
Expand All @@ -28,18 +46,31 @@ describe("getSkillPickerItems", () => {
]);
});

it("does not list skills only for the developer", () => {
const items = getSkillPickerItems([
it("lists only the skills for the given agent", () => {
const resources = [
skillResource("rill-model", ["developer"]),
skillResource("monthly-close", ["analyst"]),
]);
expect(items.map((i) => i.context.skill)).toEqual(["monthly-close"]);
skillResource("glossary", ["analyst", "developer"]),
];
expect(
getSkillPickerItems(resources, "analyst").map((i) => i.context.skill),
).toEqual(["monthly-close", "glossary"]);
expect(
getSkillPickerItems(resources, "developer").map((i) => i.context.skill),
).toEqual(["rill-model", "glossary"]);
});

it("does not list skills with a reconcile error", () => {
const items = getSkillPickerItems([
skillResource("monthly-close", ["analyst"], 'metrics view "x" not found'),
]);
const items = getSkillPickerItems(
[
skillResource(
"monthly-close",
["analyst"],
'metrics view "x" not found',
),
],
"analyst",
);
expect(items).toEqual([]);
});
});
34 changes: 26 additions & 8 deletions web-common/src/features/chat/core/context/picker/data/skills.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,35 +13,53 @@ import {
InlineContextType,
} from "@rilldata/web-common/features/chat/core/context/inline-context.ts";
import type { PickerItem } from "@rilldata/web-common/features/chat/core/context/picker/picker-tree.ts";
import { ToolName } from "@rilldata/web-common/features/chat/core/types.ts";

// Value in a skill's `agents` for the analyst agent, which answers the chat.
const ANALYST_SKILL_AGENT = "analyst";
/**
* The value in a skill's `agents` for the agent that answers a chat, by the chat's agent tool.
* Matches the agents the runtime parser accepts in a skill (see `runtime/parser/parse_skill.go`).
*/
const SKILL_AGENT_BY_CHAT_AGENT: Record<string, string> = {
[ToolName.ANALYST_AGENT]: "analyst",
[ToolName.DEVELOPER_AGENT]: "developer",
};

/**
* Returns the value in a skill's `agents` for the agent that answers a chat, or undefined if that agent has no skills.
*/
export function getSkillAgent(chatAgent: string): string | undefined {
return SKILL_AGENT_BY_CHAT_AGENT[chatAgent];
}

/**
* Creates a store that contains a flat list of the project's skills for the analyst agent.
* Creates a store that contains a flat list of the project's skills for the given skill agent.
*/
export function getSkillsPickerOptions(
client: RuntimeClient,
skillAgent: string,
): Readable<PickerItem[]> {
const skillResourcesQuery = createQuery(
getClientFilteredResourcesQueryOptions(client, ResourceKind.Skill),
queryClient,
);

return derived(skillResourcesQuery, (skillResourcesResp) =>
getSkillPickerItems(skillResourcesResp.data ?? []),
getSkillPickerItems(skillResourcesResp.data ?? [], skillAgent),
);
}

/**
* Returns picker items for the skills the analyst agent can load:
* the ones that apply to the analyst and reconciled without errors.
* Returns picker items for the skills the given agent can load:
* the ones that apply to the agent and reconciled without errors.
*/
export function getSkillPickerItems(resources: V1Resource[]): PickerItem[] {
export function getSkillPickerItems(
resources: V1Resource[],
skillAgent: string,
): PickerItem[] {
return resources
.filter(
(res) =>
res.skill?.spec?.agents?.includes(ANALYST_SKILL_AGENT) &&
res.skill?.spec?.agents?.includes(skillAgent) &&
!res.meta?.reconcileError,
)
.map((res) => {
Expand Down
16 changes: 11 additions & 5 deletions web-common/src/features/chat/core/input/ChatInput.svelte
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
<script lang="ts">
import { getEditorPlugins } from "@rilldata/web-common/features/chat/core/context/editor-plugins.svelte.ts";
import { getSkillsPickerOptions } from "@rilldata/web-common/features/chat/core/context/picker/data/skills.ts";
import {
getSkillAgent,
getSkillsPickerOptions,
} from "@rilldata/web-common/features/chat/core/context/picker/data/skills.ts";
import { useRuntimeClient } from "@rilldata/web-common/runtime-client/v2";
import { chatMounted } from "@rilldata/web-common/features/chat/layouts/sidebar/sidebar-store.ts";
import { eventBus } from "@rilldata/web-common/lib/event-bus/event-bus.ts";
Expand All @@ -26,9 +29,10 @@

let value = "";

const skillsEnabled = !!config.skills;
const skillsStore = skillsEnabled
? getSkillsPickerOptions(useRuntimeClient())
// The project's skills for this chat's agent can be picked with "/".
const skillAgent = getSkillAgent(config.agent);
const skillsStore = skillAgent
? getSkillsPickerOptions(useRuntimeClient(), skillAgent)
: readable([]);
$: hasSkills = $skillsStore.length > 0;

Expand Down Expand Up @@ -104,7 +108,9 @@
extensions: getEditorPlugins({
placeholder,
onSubmit: () => void sendMessage(),
skillOptions: skillsEnabled ? getSkillsPickerOptions : undefined,
skillOptions: skillAgent
? (client) => getSkillsPickerOptions(client, skillAgent)
: undefined,
}),
content: "",
editorProps: {
Expand Down
2 changes: 0 additions & 2 deletions web-common/src/features/chat/core/types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,4 @@ export type ChatConfig = {
emptyChatLabel: string;
placeholder: string;
minChatHeight: string;
// The project's skills for this chat's agent can be picked with "/".
skills?: boolean;
};
1 change: 0 additions & 1 deletion web-common/src/features/project/chat-context.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,5 +15,4 @@ export const projectChat = {
return m.chat_placeholder_analyst();
},
minChatHeight: "min-h-[2.5rem]",
skills: true,
} satisfies ChatConfig;
Loading