Enhance assistant prompt handling with new preset and global prompt controls

- Introduced prompt preset selection and global prompt disabling features in the assistant's message building process.
- Updated the `buildSystemPrompts` method to prioritize prompt presets and manage global prompt settings based on context and hook responses.
- Added new fields to the `HookCreateResponse` struct for prompt configuration, allowing for dynamic adjustments during message processing.
- Refactored tests to validate the integration of prompt presets and global prompt controls, ensuring robust functionality and context-aware behavior.
This commit is contained in:
Max 2025-12-02 16:24:22 +08:00
parent 42e920eed6
commit ee44432e69
3 changed files with 575 additions and 7 deletions

View file

@ -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)

View file

@ -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")
})
}

View file

@ -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