diff --git a/agent/assistant/build.go b/agent/assistant/build.go index 41131d3f..45adbce4 100644 --- a/agent/assistant/build.go +++ b/agent/assistant/build.go @@ -51,7 +51,7 @@ func (ast *Assistant) buildMessages(ctx *context.Context, messages []context.Mes } // Build and prepend system prompts (global + assistant prompts) - promptMessages := ast.buildSystemPrompts(ctx) + promptMessages := ast.buildSystemPrompts(ctx, createResponse) if len(promptMessages) > 0 { finalMessages = append(promptMessages, finalMessages...) } @@ -60,25 +60,41 @@ func (ast *Assistant) buildMessages(ctx *context.Context, messages []context.Mes } // buildSystemPrompts builds system prompt messages from global prompts and assistant prompts -// Order: Global prompts (if not disabled) -> Assistant prompts +// Order: Global prompts (if not disabled) -> Assistant prompts (or preset) // Variables are parsed with context information -func (ast *Assistant) buildSystemPrompts(ctx *context.Context) []context.Message { +// +// Priority for prompt preset selection: +// 1. createResponse.PromptPreset (highest) +// 2. ctx.Metadata["__prompt_preset"] +// 3. ast.Prompts (default) +// +// Priority for disable global prompts: +// 1. createResponse.DisableGlobalPrompts (highest) +// 2. ctx.Metadata["__disable_global_prompts"] +// 3. ast.DisableGlobalPrompts (default) +func (ast *Assistant) buildSystemPrompts(ctx *context.Context, createResponse *context.HookCreateResponse) []context.Message { // Build context variables from ctx and ast ctxVars := ast.buildContextVariables(ctx) + // Determine if global prompts should be disabled + disableGlobal := ast.shouldDisableGlobalPrompts(ctx, createResponse) + + // Get assistant prompts (default or preset) + assistantPrompts := ast.getAssistantPrompts(ctx, createResponse) + var allPrompts []store.Prompt // 1. Add global prompts (if not disabled) - if !ast.DisableGlobalPrompts && len(globalPrompts) > 0 { + if !disableGlobal && 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 { + // 2. Add assistant prompts (default or preset) + if len(assistantPrompts) > 0 { // Parse assistant prompts with context variables - parsedAssistant := store.Prompts(ast.Prompts).Parse(ctxVars) + parsedAssistant := store.Prompts(assistantPrompts).Parse(ctxVars) allPrompts = append(allPrompts, parsedAssistant...) } @@ -103,6 +119,61 @@ func (ast *Assistant) buildSystemPrompts(ctx *context.Context) []context.Message return messages } +// shouldDisableGlobalPrompts determines if global prompts should be disabled +// Priority: createResponse > ctx.Metadata > ast.DisableGlobalPrompts +func (ast *Assistant) shouldDisableGlobalPrompts(ctx *context.Context, createResponse *context.HookCreateResponse) bool { + // Priority 1: Hook response (highest) + if createResponse != nil && createResponse.DisableGlobalPrompts != nil { + return *createResponse.DisableGlobalPrompts + } + + // Priority 2: ctx.Metadata["__disable_global_prompts"] + if ctx != nil && ctx.Metadata != nil { + if disable, ok := ctx.Metadata["__disable_global_prompts"].(bool); ok { + return disable + } + } + + // Priority 3: Assistant configuration (default) + return ast.DisableGlobalPrompts +} + +// getAssistantPrompts returns the assistant prompts based on preset selection +// Priority: createResponse.PromptPreset > ctx.Metadata["__prompt_preset"] > ast.Prompts +func (ast *Assistant) getAssistantPrompts(ctx *context.Context, createResponse *context.HookCreateResponse) []store.Prompt { + // Get preset key + presetKey := ast.getPromptPresetKey(ctx, createResponse) + + // If preset key is specified and exists, use it + if presetKey != "" && ast.PromptPresets != nil { + if presets, ok := ast.PromptPresets[presetKey]; ok && len(presets) > 0 { + return presets + } + } + + // Fallback to default prompts + return ast.Prompts +} + +// getPromptPresetKey returns the prompt preset key +// Priority: createResponse.PromptPreset > ctx.Metadata["__prompt_preset"] +func (ast *Assistant) getPromptPresetKey(ctx *context.Context, createResponse *context.HookCreateResponse) string { + // Priority 1: Hook response (highest) + if createResponse != nil && createResponse.PromptPreset != "" { + return createResponse.PromptPreset + } + + // Priority 2: ctx.Metadata["__prompt_preset"] + if ctx != nil && ctx.Metadata != nil { + if preset, ok := ctx.Metadata["__prompt_preset"].(string); ok && preset != "" { + return preset + } + } + + // No preset specified + return "" +} + // 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) diff --git a/agent/assistant/build_prompts_test.go b/agent/assistant/build_prompts_test.go index a8f029be..a96fda4c 100644 --- a/agent/assistant/build_prompts_test.go +++ b/agent/assistant/build_prompts_test.go @@ -1,6 +1,8 @@ package assistant_test import ( + stdContext "context" + "strings" "testing" "github.com/stretchr/testify/assert" @@ -12,6 +14,43 @@ import ( "github.com/yaoapp/yao/openapi/oauth/types" ) +// containsString is a helper to check if a content (string or interface{}) contains a substring +func containsString(content interface{}, substr string) bool { + switch v := content.(type) { + case string: + return strings.Contains(v, substr) + default: + return false + } +} + +// newPromptTestContext creates a context suitable for prompt testing with Create Hook +func newPromptTestContext(chatID, assistantID string) *context.Context { + return &context.Context{ + Context: stdContext.Background(), + ChatID: chatID, + AssistantID: assistantID, + Locale: "en-us", + Theme: "light", + Client: context.Client{ + Type: "web", + UserAgent: "TestAgent/1.0", + IP: "127.0.0.1", + }, + Referer: context.RefererAPI, + Accept: context.AcceptWebCUI, + Metadata: make(map[string]interface{}), + Authorized: &types.AuthorizedInfo{ + Subject: "test-user", + ClientID: "test-client-id", + UserID: "test-user-123", + TeamID: "test-team-456", + TenantID: "test-tenant-789", + SessionID: "test-session-id", + }, + } +} + func TestBuildSystemPromptsIntegration(t *testing.T) { testutils.Prepare(t) defer testutils.Clean(t) @@ -344,4 +383,458 @@ func TestBuildSystemPromptsIntegration(t *testing.T) { } assert.True(t, found, "Should find global prompt with all variable types replaced") }) + + t.Run("PromptPresetFromHook", func(t *testing.T) { + // Load fullfields assistant which has prompt_presets + ast, err := assistant.Get("tests.fullfields") + require.NoError(t, err) + require.NotNil(t, ast.PromptPresets) + require.Contains(t, ast.PromptPresets, "chat.friendly") + + ctx := &context.Context{} + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test preset from hook"}, + } + + // Hook returns prompt_preset + createResponse := &context.HookCreateResponse{ + PromptPreset: "chat.friendly", + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should have system prompts from the preset + hasSystemPrompt := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + hasSystemPrompt = true + // Verify it's from the friendly preset (check content) + assert.Contains(t, msg.Content, "friendly", "Should use friendly preset prompts") + break + } + } + assert.True(t, hasSystemPrompt, "Should have system prompts from preset") + }) + + t.Run("PromptPresetFromMetadata", func(t *testing.T) { + // Load fullfields assistant which has prompt_presets + ast, err := assistant.Get("tests.fullfields") + require.NoError(t, err) + + ctx := &context.Context{ + Metadata: map[string]interface{}{ + "__prompt_preset": "chat.professional", + }, + } + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test preset from metadata"}, + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, nil) + require.NoError(t, err) + + // Should have system prompts from the preset + hasSystemPrompt := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + hasSystemPrompt = true + // Verify it's from the professional preset + assert.Contains(t, msg.Content, "professional", "Should use professional preset prompts") + break + } + } + assert.True(t, hasSystemPrompt, "Should have system prompts from preset") + }) + + t.Run("PromptPresetHookOverridesMetadata", func(t *testing.T) { + // Load fullfields assistant + ast, err := assistant.Get("tests.fullfields") + require.NoError(t, err) + + ctx := &context.Context{ + Metadata: map[string]interface{}{ + "__prompt_preset": "chat.professional", // Lower priority + }, + } + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test hook overrides metadata"}, + } + + // Hook returns different preset (higher priority) + createResponse := &context.HookCreateResponse{ + PromptPreset: "chat.friendly", + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should use hook's preset, not metadata's + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + assert.Contains(t, msg.Content, "friendly", "Hook preset should override metadata preset") + break + } + } + }) + + t.Run("PromptPresetNotFound", func(t *testing.T) { + // Load fullfields assistant + ast, err := assistant.Get("tests.fullfields") + require.NoError(t, err) + + ctx := &context.Context{ + Metadata: map[string]interface{}{ + "__prompt_preset": "non.existent.preset", + }, + } + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test non-existent preset"}, + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, nil) + require.NoError(t, err) + + // Should fallback to default prompts (not crash) + hasSystemPrompt := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + hasSystemPrompt = true + break + } + } + assert.True(t, hasSystemPrompt, "Should fallback to default prompts when preset not found") + }) + + t.Run("DisableGlobalPromptsFromHook", func(t *testing.T) { + // Set global prompts + assistant.SetGlobalPrompts([]store.Prompt{ + {Role: "system", Content: "GLOBAL_PROMPT_MARKER"}, + }) + defer assistant.SetGlobalPrompts(nil) + + // Load an assistant that does NOT disable global prompts + ast, err := assistant.Get("yaobots") + require.NoError(t, err) + require.False(t, ast.DisableGlobalPrompts) + + ctx := &context.Context{} + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test disable from hook"}, + } + + // Hook disables global prompts + disableTrue := true + createResponse := &context.HookCreateResponse{ + DisableGlobalPrompts: &disableTrue, + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should NOT have global prompt + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + assert.NotContains(t, msg.Content, "GLOBAL_PROMPT_MARKER", "Global prompts should be disabled by hook") + } + } + }) + + t.Run("DisableGlobalPromptsFromMetadata", func(t *testing.T) { + // Set global prompts + assistant.SetGlobalPrompts([]store.Prompt{ + {Role: "system", Content: "GLOBAL_PROMPT_MARKER_2"}, + }) + defer assistant.SetGlobalPrompts(nil) + + // Load an assistant that does NOT disable global prompts + ast, err := assistant.Get("yaobots") + require.NoError(t, err) + + ctx := &context.Context{ + Metadata: map[string]interface{}{ + "__disable_global_prompts": true, + }, + } + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test disable from metadata"}, + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, nil) + require.NoError(t, err) + + // Should NOT have global prompt + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + assert.NotContains(t, msg.Content, "GLOBAL_PROMPT_MARKER_2", "Global prompts should be disabled by metadata") + } + } + }) + + t.Run("EnableGlobalPromptsOverrideAssistant", func(t *testing.T) { + // Set global prompts + assistant.SetGlobalPrompts([]store.Prompt{ + {Role: "system", Content: "GLOBAL_ENABLED_MARKER"}, + }) + defer assistant.SetGlobalPrompts(nil) + + // 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{} + + messages := []context.Message{ + {Role: context.RoleUser, Content: "Test enable override"}, + } + + // Hook enables global prompts (overrides assistant's disable) + disableFalse := false + createResponse := &context.HookCreateResponse{ + DisableGlobalPrompts: &disableFalse, + } + + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should have global prompt (hook enabled it) + found := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem && msg.Content == "GLOBAL_ENABLED_MARKER" { + found = true + break + } + } + assert.True(t, found, "Global prompts should be enabled by hook override") + }) +} + +// TestPromptPresetAssistant tests the tests.promptpreset assistant with Create Hook +func TestPromptPresetAssistant(t *testing.T) { + testutils.Prepare(t) + defer testutils.Clean(t) + + t.Run("LoadPromptPresetAssistant", func(t *testing.T) { + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + require.NotNil(t, ast) + + assert.Equal(t, "tests.promptpreset", ast.ID) + assert.Equal(t, "Prompt Preset Test", ast.Name) + assert.False(t, ast.DisableGlobalPrompts) + + // Should have prompt presets loaded + require.NotNil(t, ast.PromptPresets) + assert.Contains(t, ast.PromptPresets, "mode.friendly") + assert.Contains(t, ast.PromptPresets, "mode.professional") + + // Should have script + assert.NotNil(t, ast.Script) + }) + + t.Run("CreateHookSelectsFriendlyPreset", func(t *testing.T) { + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-friendly-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "use friendly mode please"}, + } + + // Call Create hook + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + require.NotNil(t, createResponse) + assert.Equal(t, "mode.friendly", createResponse.PromptPreset) + + // Build request + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should have friendly preset marker in one of the system messages + found := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem && containsString(msg.Content, "FRIENDLY_PRESET_MARKER") { + found = true + break + } + } + assert.True(t, found, "Should use friendly preset from Create Hook") + }) + + t.Run("CreateHookSelectsProfessionalPreset", func(t *testing.T) { + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-professional-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "use professional tone"}, + } + + // Call Create hook + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + require.NotNil(t, createResponse) + assert.Equal(t, "mode.professional", createResponse.PromptPreset) + + // Build request + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should have professional preset marker in one of the system messages + found := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem && containsString(msg.Content, "PROFESSIONAL_PRESET_MARKER") { + found = true + break + } + } + assert.True(t, found, "Should use professional preset from Create Hook") + }) + + t.Run("CreateHookDisablesGlobalPrompts", func(t *testing.T) { + // Set global prompts + assistant.SetGlobalPrompts([]store.Prompt{ + {Role: "system", Content: "GLOBAL_MARKER_FOR_DISABLE_TEST"}, + }) + defer assistant.SetGlobalPrompts(nil) + + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-disable-global-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "disable global prompts"}, + } + + // Call Create hook + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + require.NotNil(t, createResponse) + require.NotNil(t, createResponse.DisableGlobalPrompts) + assert.True(t, *createResponse.DisableGlobalPrompts) + + // Build request + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should NOT have global prompt + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + assert.NotContains(t, msg.Content, "GLOBAL_MARKER_FOR_DISABLE_TEST") + } + } + }) + + t.Run("CreateHookPresetAndDisableGlobal", func(t *testing.T) { + // Set global prompts + assistant.SetGlobalPrompts([]store.Prompt{ + {Role: "system", Content: "GLOBAL_MARKER_COMBINED_TEST"}, + }) + defer assistant.SetGlobalPrompts(nil) + + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-combined-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "friendly no global"}, + } + + // Call Create hook + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + require.NotNil(t, createResponse) + assert.Equal(t, "mode.friendly", createResponse.PromptPreset) + require.NotNil(t, createResponse.DisableGlobalPrompts) + assert.True(t, *createResponse.DisableGlobalPrompts) + + // Build request + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should have friendly preset but NOT global + hasFriendly := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem { + assert.NotContains(t, msg.Content, "GLOBAL_MARKER_COMBINED_TEST") + if containsString(msg.Content, "FRIENDLY_PRESET_MARKER") { + hasFriendly = true + } + } + } + assert.True(t, hasFriendly, "Should have friendly preset") + }) + + t.Run("CreateHookUnknownPresetFallback", func(t *testing.T) { + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-unknown-preset-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "unknown preset test"}, + } + + // Call Create hook + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + require.NotNil(t, createResponse) + assert.Equal(t, "non.existent.preset", createResponse.PromptPreset) + + // Build request - should not error, fallback to default + finalMessages, _, err := ast.BuildRequest(ctx, messages, createResponse) + require.NoError(t, err) + + // Should fallback to default prompts + found := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem && containsString(msg.Content, "DEFAULT_PROMPT_MARKER") { + found = true + break + } + } + assert.True(t, found, "Should fallback to default prompts when preset not found") + }) + + t.Run("CreateHookReturnsNull", func(t *testing.T) { + ast, err := assistant.Get("tests.promptpreset") + require.NoError(t, err) + + ctx := newPromptTestContext("chat-null-test", "tests.promptpreset") + + messages := []context.Message{ + {Role: context.RoleUser, Content: "just a normal message"}, + } + + // Call Create hook - should return nil + createResponse, err := ast.Script.Create(ctx, messages) + require.NoError(t, err) + assert.Nil(t, createResponse) + + // Build request with nil createResponse + finalMessages, _, err := ast.BuildRequest(ctx, messages, nil) + require.NoError(t, err) + + // Should use default prompts + found := false + for _, msg := range finalMessages { + if msg.Role == context.RoleSystem && containsString(msg.Content, "DEFAULT_PROMPT_MARKER") { + found = true + break + } + } + assert.True(t, found, "Should use default prompts when hook returns null") + }) } diff --git a/agent/context/types.go b/agent/context/types.go index 4a5f6587..d62cc1e5 100644 --- a/agent/context/types.go +++ b/agent/context/types.go @@ -326,6 +326,10 @@ type HookCreateResponse struct { // MCP configuration - allow hook to add/override MCP servers for this request MCPServers []MCPServerConfig `json:"mcp_servers,omitempty"` + // Prompt configuration + PromptPreset string `json:"prompt_preset,omitempty"` // Select prompt preset (e.g., "chat.friendly", "task.analysis") + DisableGlobalPrompts *bool `json:"disable_global_prompts,omitempty"` // Temporarily disable global prompts for this request + // Context adjustments - allow hook to modify context fields AssistantID string `json:"assistant_id,omitempty"` // Override assistant ID Connector string `json:"connector,omitempty"` // Override connector