Implement global prompts functionality in the assistant module

- Added methods to set and retrieve global prompts, enhancing the assistant's capabilities.
- Updated the assistant's message building process to include global prompts, ensuring context-aware parsing.
- Introduced tests to validate the integration and functionality of global prompts within the assistant.
- Improved context variable handling for prompt parsing, supporting dynamic content generation.
This commit is contained in:
Max 2025-12-02 16:03:31 +08:00
parent 45029269af
commit 42e920eed6
5 changed files with 540 additions and 19 deletions

View file

@ -3,8 +3,10 @@ package assistant
import (
"fmt"
"github.com/spf13/cast"
"github.com/yaoapp/gou/json"
"github.com/yaoapp/yao/agent/context"
store "github.com/yaoapp/yao/agent/store/types"
)
// BuildRequest build the LLM request
@ -48,29 +50,146 @@ func (ast *Assistant) buildMessages(ctx *context.Context, messages []context.Mes
finalMessages = append([]context.Message{mcpSamplesMsg}, finalMessages...)
}
// ⚠️ Just for testing, will remove later
// If we have prompts, prepend them to the beginning
if len(ast.Prompts) > 0 {
promptMessages := make([]context.Message, 0, len(ast.Prompts))
for _, prompt := range ast.Prompts {
msg := context.Message{
Role: context.MessageRole(prompt.Role),
Content: prompt.Content,
}
// Add name if provided
if prompt.Name != "" {
name := prompt.Name
msg.Name = &name
}
promptMessages = append(promptMessages, msg)
}
// Prepend prompt messages to the beginning
// Build and prepend system prompts (global + assistant prompts)
promptMessages := ast.buildSystemPrompts(ctx)
if len(promptMessages) > 0 {
finalMessages = append(promptMessages, finalMessages...)
}
return finalMessages, nil
}
// buildSystemPrompts builds system prompt messages from global prompts and assistant prompts
// Order: Global prompts (if not disabled) -> Assistant prompts
// Variables are parsed with context information
func (ast *Assistant) buildSystemPrompts(ctx *context.Context) []context.Message {
// Build context variables from ctx and ast
ctxVars := ast.buildContextVariables(ctx)
var allPrompts []store.Prompt
// 1. Add global prompts (if not disabled)
if !ast.DisableGlobalPrompts && len(globalPrompts) > 0 {
// Parse global prompts with context variables
parsedGlobal := store.Prompts(globalPrompts).Parse(ctxVars)
allPrompts = append(allPrompts, parsedGlobal...)
}
// 2. Add assistant prompts
if len(ast.Prompts) > 0 {
// Parse assistant prompts with context variables
parsedAssistant := store.Prompts(ast.Prompts).Parse(ctxVars)
allPrompts = append(allPrompts, parsedAssistant...)
}
// Convert to context.Message slice
if len(allPrompts) == 0 {
return nil
}
messages := make([]context.Message, 0, len(allPrompts))
for _, prompt := range allPrompts {
msg := context.Message{
Role: context.MessageRole(prompt.Role),
Content: prompt.Content,
}
if prompt.Name != "" {
name := prompt.Name
msg.Name = &name
}
messages = append(messages, msg)
}
return messages
}
// buildContextVariables extracts context variables from Context and Assistant for prompt parsing
func (ast *Assistant) buildContextVariables(ctx *context.Context) map[string]string {
vars := make(map[string]string)
// Get locale from ctx (default to empty)
locale := ""
if ctx != nil && ctx.Locale != "" {
locale = ctx.Locale
}
// Assistant info (with locale support)
if ast != nil {
if ast.ID != "" {
vars["ASSISTANT_ID"] = ast.ID
}
// Use localized name and description
name := ast.GetName(locale)
if name != "" {
vars["ASSISTANT_NAME"] = name
}
description := ast.GetDescription(locale)
if description != "" {
vars["ASSISTANT_DESCRIPTION"] = description
}
if ast.Type != "" {
vars["ASSISTANT_TYPE"] = ast.Type
}
}
if ctx == nil {
return vars
}
// Basic context info
if ctx.ChatID != "" {
vars["CHAT_ID"] = ctx.ChatID
}
if ctx.Locale != "" {
vars["LOCALE"] = ctx.Locale
}
if ctx.Theme != "" {
vars["THEME"] = ctx.Theme
}
if ctx.Route != "" {
vars["ROUTE"] = ctx.Route
}
if ctx.Referer != "" {
vars["REFERER"] = ctx.Referer
}
// Client info (only non-sensitive fields)
if ctx.Client.Type != "" {
vars["CLIENT_TYPE"] = ctx.Client.Type
}
// Authorized info (only internal IDs, no PII)
// Note: USER_SUBJECT and CLIENT_IP are excluded for privacy/GDPR compliance
if ctx.Authorized != nil {
if ctx.Authorized.UserID != "" {
vars["USER_ID"] = ctx.Authorized.UserID
}
if ctx.Authorized.TeamID != "" {
vars["TEAM_ID"] = ctx.Authorized.TeamID
}
if ctx.Authorized.TenantID != "" {
vars["TENANT_ID"] = ctx.Authorized.TenantID
}
}
// Metadata - custom variables from ctx.Metadata
// All metadata keys are exposed as $CTX.{KEY}
// Supports string, int, uint, float, bool types
if ctx.Metadata != nil {
for key, value := range ctx.Metadata {
if value == nil {
continue
}
strVal := cast.ToString(value)
if strVal != "" {
vars[key] = strVal
}
}
}
return vars
}
// buildCompletionOptions builds completion options from multiple sources
// Priority (lowest to highest, later overrides earlier): ast > ctx > createResponse
// The priority means: if createResponse has a value, use it; else use ctx; else use ast

View file

@ -0,0 +1,347 @@
package assistant_test
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/yao/agent/assistant"
"github.com/yaoapp/yao/agent/context"
store "github.com/yaoapp/yao/agent/store/types"
"github.com/yaoapp/yao/agent/testutils"
"github.com/yaoapp/yao/openapi/oauth/types"
)
func TestBuildSystemPromptsIntegration(t *testing.T) {
testutils.Prepare(t)
defer testutils.Clean(t)
t.Run("AssistantWithLocale", func(t *testing.T) {
// Load an assistant with locales
ast, err := assistant.Get("tests.fullfields")
require.NoError(t, err)
ctx := &context.Context{
Locale: "zh-cn",
Authorized: &types.AuthorizedInfo{
UserID: "test-user-123",
TeamID: "test-team-456",
},
Metadata: map[string]interface{}{
"CUSTOM_VAR": "custom-value",
"INT_VAR": 42,
"BOOL_VAR": true,
},
}
// Build request to test the full flow
messages := []context.Message{
{Role: context.RoleUser, Content: "Hello"},
}
finalMessages, options, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
require.NotNil(t, options)
// Should have system prompts prepended
assert.Greater(t, len(finalMessages), 1)
// First messages should be system prompts
hasSystemPrompt := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem {
hasSystemPrompt = true
break
}
}
assert.True(t, hasSystemPrompt, "Should have system prompts")
})
t.Run("DisableGlobalPrompts", func(t *testing.T) {
// Load fullfields assistant which has disable_global_prompts: true
ast, err := assistant.Get("tests.fullfields")
require.NoError(t, err)
require.True(t, ast.DisableGlobalPrompts)
ctx := &context.Context{
Locale: "en-us",
}
messages := []context.Message{
{Role: context.RoleUser, Content: "Hello"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Should still have assistant prompts
hasSystemPrompt := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem {
hasSystemPrompt = true
break
}
}
assert.True(t, hasSystemPrompt, "Should have assistant prompts even with global disabled")
})
t.Run("MetadataTypeConversion", func(t *testing.T) {
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{
Metadata: map[string]interface{}{
"STRING_VAL": "hello",
"INT_VAL": 123,
"INT64_VAL": int64(456),
"FLOAT_VAL": 3.14,
"BOOL_TRUE": true,
"BOOL_FALSE": false,
"UINT_VAL": uint(789),
"NIL_VAL": nil,
"EMPTY_VAL": "",
"ZERO_INT": 0,
"ZERO_FLOAT": 0.0,
},
}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test metadata"},
}
// This should not panic
_, _, err = ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
})
t.Run("AuthorizedInfoPrivacy", func(t *testing.T) {
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{
Authorized: &types.AuthorizedInfo{
UserID: "user-123",
Subject: "user@example.com", // PII - should not be exposed
TeamID: "team-456",
TenantID: "tenant-789",
},
Client: context.Client{
Type: "web",
IP: "192.168.1.1", // Should not be exposed
},
}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test privacy"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Check that sensitive info is not in any system prompts
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem {
assert.NotContains(t, msg.Content, "user@example.com", "Subject should not be in prompts")
assert.NotContains(t, msg.Content, "192.168.1.1", "IP should not be in prompts")
}
}
})
t.Run("ContextVariablesInPrompts", func(t *testing.T) {
// Set up global prompts with variables
assistant.SetGlobalPrompts([]store.Prompt{
{Role: "system", Content: "User ID: $CTX.USER_ID, Team: $CTX.TEAM_ID, Custom: $CTX.MY_VAR"},
})
defer assistant.SetGlobalPrompts(nil)
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{
Authorized: &types.AuthorizedInfo{
UserID: "user-abc",
TeamID: "team-xyz",
},
Metadata: map[string]interface{}{
"MY_VAR": "my-value",
},
}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test variables"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Find the global prompt and verify variables are replaced
found := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem && !found {
if assert.Contains(t, msg.Content, "User ID: user-abc") {
found = true
assert.Contains(t, msg.Content, "Team: team-xyz")
assert.Contains(t, msg.Content, "Custom: my-value")
}
}
}
assert.True(t, found, "Should find global prompt with replaced variables")
})
t.Run("SystemVariablesReplacement", func(t *testing.T) {
// Set up global prompts with $SYS.* variables
assistant.SetGlobalPrompts([]store.Prompt{
{Role: "system", Content: "Time: $SYS.TIME, Date: $SYS.DATE, Datetime: $SYS.DATETIME, Weekday: $SYS.WEEKDAY"},
})
defer assistant.SetGlobalPrompts(nil)
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test system variables"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Find the global prompt and verify $SYS.* variables are replaced
found := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem {
// Should NOT contain $SYS. prefix (variables should be replaced)
if !assert.NotContains(t, msg.Content, "$SYS.TIME") {
continue
}
if !assert.NotContains(t, msg.Content, "$SYS.DATE") {
continue
}
if !assert.NotContains(t, msg.Content, "$SYS.DATETIME") {
continue
}
if !assert.NotContains(t, msg.Content, "$SYS.WEEKDAY") {
continue
}
// Should contain "Time:", "Date:", etc. with actual values
assert.Contains(t, msg.Content, "Time:")
assert.Contains(t, msg.Content, "Date:")
assert.Contains(t, msg.Content, "Datetime:")
assert.Contains(t, msg.Content, "Weekday:")
found = true
break
}
}
assert.True(t, found, "Should find global prompt with replaced $SYS.* variables")
})
t.Run("EnvVariablesReplacement", func(t *testing.T) {
// Set test environment variable
t.Setenv("TEST_PROMPT_VAR", "env-test-value")
// Set up global prompts with $ENV.* variables
assistant.SetGlobalPrompts([]store.Prompt{
{Role: "system", Content: "Env Value: $ENV.TEST_PROMPT_VAR, Not Exist: $ENV.NOT_EXIST_VAR_XYZ"},
})
defer assistant.SetGlobalPrompts(nil)
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test env variables"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Find the global prompt and verify $ENV.* variables are replaced
found := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem {
// Should NOT contain $ENV. prefix for existing vars
if !assert.NotContains(t, msg.Content, "$ENV.TEST_PROMPT_VAR") {
continue
}
// Should contain the actual env value
assert.Contains(t, msg.Content, "Env Value: env-test-value")
// Non-existent env var should be replaced with empty string
assert.Contains(t, msg.Content, "Not Exist: ")
assert.NotContains(t, msg.Content, "$ENV.NOT_EXIST_VAR_XYZ")
found = true
break
}
}
assert.True(t, found, "Should find global prompt with replaced $ENV.* variables")
})
t.Run("AllVariableTypesReplacement", func(t *testing.T) {
// Set test environment variable
t.Setenv("TEST_APP_NAME", "MyTestApp")
// Set up global prompts with all variable types
assistant.SetGlobalPrompts([]store.Prompt{
{Role: "system", Content: `System Info:
- Time: $SYS.TIME
- Date: $SYS.DATE
- App: $ENV.TEST_APP_NAME
- User: $CTX.USER_ID
- Custom: $CTX.CUSTOM_KEY
- Assistant: $CTX.ASSISTANT_NAME`},
})
defer assistant.SetGlobalPrompts(nil)
ast, err := assistant.Get("yaobots")
require.NoError(t, err)
ctx := &context.Context{
Authorized: &types.AuthorizedInfo{
UserID: "all-vars-user",
},
Metadata: map[string]interface{}{
"CUSTOM_KEY": "custom-value-123",
},
}
messages := []context.Message{
{Role: context.RoleUser, Content: "Test all variables"},
}
finalMessages, _, err := ast.BuildRequest(ctx, messages, nil)
require.NoError(t, err)
// Find the global prompt and verify ALL variable types are replaced
found := false
for _, msg := range finalMessages {
if msg.Role == context.RoleSystem && !found {
content := msg.Content
// Check $SYS.* replaced
if assert.NotContains(t, content, "$SYS.TIME") &&
assert.NotContains(t, content, "$SYS.DATE") {
// Check $ENV.* replaced
assert.NotContains(t, content, "$ENV.TEST_APP_NAME")
assert.Contains(t, content, "App: MyTestApp")
// Check $CTX.* replaced
assert.NotContains(t, content, "$CTX.USER_ID")
assert.Contains(t, content, "User: all-vars-user")
assert.NotContains(t, content, "$CTX.CUSTOM_KEY")
assert.Contains(t, content, "Custom: custom-value-123")
// Check assistant name from $CTX.ASSISTANT_NAME
assert.NotContains(t, content, "$CTX.ASSISTANT_NAME")
found = true
}
}
}
assert.True(t, found, "Should find global prompt with all variable types replaced")
})
}

View file

@ -27,8 +27,9 @@ var loaded = NewCache(200) // 200 is the default capacity
var storage store.Store = nil
var search interface{} = nil
var modelCapabilities map[string]gouOpenAI.Capabilities = map[string]gouOpenAI.Capabilities{}
var defaultConnector string = "" // default connector
var globalUses *context.Uses = nil // global uses configuration from agent.yml
var defaultConnector string = "" // default connector
var globalUses *context.Uses = nil // global uses configuration from agent.yml
var globalPrompts []store.Prompt = nil // global prompts from agent/prompts.yml
// LoadBuiltIn load the built-in assistants
func LoadBuiltIn() error {
@ -145,6 +146,20 @@ func SetGlobalUses(uses *context.Uses) {
globalUses = uses
}
// SetGlobalPrompts set the global prompts from agent/prompts.yml
func SetGlobalPrompts(prompts []store.Prompt) {
globalPrompts = prompts
}
// GetGlobalPrompts returns the global prompts with variables parsed
// ctx: context variables for parsing $CTX.* variables
func GetGlobalPrompts(ctx map[string]string) []store.Prompt {
if len(globalPrompts) == 0 {
return nil
}
return store.Prompts(globalPrompts).Parse(ctx)
}
// SetCache set the cache
func SetCache(capacity int) {
ClearCache()

View file

@ -200,6 +200,11 @@ func initAssistant() error {
assistant.SetGlobalUses(globalUses)
}
// Set global prompts
if len(agentDSL.GlobalPrompts) > 0 {
assistant.SetGlobalPrompts(agentDSL.GlobalPrompts)
}
if agentDSL.Models != nil {
assistant.SetModelCapabilities(agentDSL.Models)
}

View file

@ -7,6 +7,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/yao/agent/assistant"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test"
)
@ -171,3 +172,37 @@ func TestGlobalPromptsContent(t *testing.T) {
"Raw prompts should contain variable placeholders")
})
}
func TestAssistantGlobalPrompts(t *testing.T) {
prepare(t)
defer test.Clean()
t.Run("AssistantModuleReceivesGlobalPrompts", func(t *testing.T) {
// Verify assistant module has global prompts
prompts := assistant.GetGlobalPrompts(nil)
require.NotNil(t, prompts)
require.Greater(t, len(prompts), 0)
// Should be parsed (no $SYS.* variables)
content := prompts[0].Content
assert.NotContains(t, content, "$SYS.DATETIME")
})
t.Run("AssistantModuleParsesWithContext", func(t *testing.T) {
ctx := map[string]string{
"USER_ID": "assistant-test-user",
"LOCALE": "en-US",
}
prompts := assistant.GetGlobalPrompts(ctx)
require.NotNil(t, prompts)
// $SYS.* should be replaced
content := prompts[0].Content
assert.NotContains(t, content, "$SYS.")
// Should contain current time info
now := time.Now()
assert.Contains(t, content, now.Format("2006-01-02"))
})
}