From 42e920eed6c08929e52f5bbdac701c9894d2ff4a Mon Sep 17 00:00:00 2001 From: Max Date: Tue, 2 Dec 2025 16:03:31 +0800 Subject: [PATCH] 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. --- agent/assistant/build.go | 153 ++++++++++-- agent/assistant/build_prompts_test.go | 347 ++++++++++++++++++++++++++ agent/assistant/load.go | 19 +- agent/load.go | 5 + agent/load_test.go | 35 +++ 5 files changed, 540 insertions(+), 19 deletions(-) create mode 100644 agent/assistant/build_prompts_test.go diff --git a/agent/assistant/build.go b/agent/assistant/build.go index d7feef05..41131d3f 100644 --- a/agent/assistant/build.go +++ b/agent/assistant/build.go @@ -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 diff --git a/agent/assistant/build_prompts_test.go b/agent/assistant/build_prompts_test.go new file mode 100644 index 00000000..a8f029be --- /dev/null +++ b/agent/assistant/build_prompts_test.go @@ -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") + }) +} diff --git a/agent/assistant/load.go b/agent/assistant/load.go index c15c228b..e62b6789 100644 --- a/agent/assistant/load.go +++ b/agent/assistant/load.go @@ -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() diff --git a/agent/load.go b/agent/load.go index 461dfe89..44fa1494 100644 --- a/agent/load.go +++ b/agent/load.go @@ -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) } diff --git a/agent/load_test.go b/agent/load_test.go index 40c95069..602cc094 100644 --- a/agent/load_test.go +++ b/agent/load_test.go @@ -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")) + }) +}