diff --git a/runtime/ai/analyst_agent.go b/runtime/ai/analyst_agent.go
index 7a2cb41aa3e0..1cd2017063f9 100644
--- a/runtime/ai/analyst_agent.go
+++ b/runtime/ai/analyst_agent.go
@@ -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
diff --git a/runtime/ai/developer_agent.go b/runtime/ai/developer_agent.go
index 85ed4a2bea70..a712507e2b56 100644
--- a/runtime/ai/developer_agent.go
+++ b/runtime/ai/developer_agent.go
@@ -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)}
diff --git a/runtime/ai/skill_references.go b/runtime/ai/skill_references.go
index 2af3709447e0..9eff8d36c9cb 100644
--- a/runtime/ai/skill_references.go
+++ b/runtime/ai/skill_references.go
@@ -1,6 +1,8 @@
package ai
import (
+ "context"
+ "errors"
"regexp"
"slices"
)
@@ -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 {
@@ -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
diff --git a/runtime/ai/skill_references_test.go b/runtime/ai/skill_references_test.go
index 08a357378b81..e63ddcaa4764 100644
--- a/runtime/ai/skill_references_test.go
+++ b/runtime/ai/skill_references_test.go
@@ -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
@@ -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.`,
},
@@ -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 := `type="skill" skill="churn-review" as a model, then ` +
+ `skill="churn-review" type="skill" again. ` +
+ `type="skill" skill="does-not-exist" ` +
+ `type="skill" skill="monthly-close" ` +
+ `type="skill" skill="dev-conventions"`
+ 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 := `type="skill" skill="churn-review" 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.")
+}
diff --git a/web-common/src/features/chat/core/context/picker/data/skills.spec.ts b/web-common/src/features/chat/core/context/picker/data/skills.spec.ts
index d5381db583d3..79f133c8e034 100644
--- a/web-common/src/features/chat/core/context/picker/data/skills.spec.ts
+++ b/web-common/src/features/chat/core/context/picker/data/skills.spec.ts
@@ -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,
@@ -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([
@@ -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([]);
});
});
diff --git a/web-common/src/features/chat/core/context/picker/data/skills.ts b/web-common/src/features/chat/core/context/picker/data/skills.ts
index 6e7ecc38b08a..2cd1e6a2dff3 100644
--- a/web-common/src/features/chat/core/context/picker/data/skills.ts
+++ b/web-common/src/features/chat/core/context/picker/data/skills.ts
@@ -13,15 +13,30 @@ 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 = {
+ [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 {
const skillResourcesQuery = createQuery(
getClientFilteredResourcesQueryOptions(client, ResourceKind.Skill),
@@ -29,19 +44,22 @@ export function getSkillsPickerOptions(
);
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) => {
diff --git a/web-common/src/features/chat/core/input/ChatInput.svelte b/web-common/src/features/chat/core/input/ChatInput.svelte
index a4e5c432aedb..2efa9d0de7d9 100644
--- a/web-common/src/features/chat/core/input/ChatInput.svelte
+++ b/web-common/src/features/chat/core/input/ChatInput.svelte
@@ -1,6 +1,9 @@