From 8dca4719c0348ce502de54599240affc5d0b207d Mon Sep 17 00:00:00 2001 From: Max Date: Tue, 11 Nov 2025 11:23:20 +0800 Subject: [PATCH] Remove deprecated agent API files and refactor agent loading logic - Deleted obsolete agent API files (agent.go, api.go, api_test.go, types.go) to streamline the codebase. - Refactored the agent loading logic to initialize the API instance correctly, ensuring proper integration with the new structure. - Updated context handling to improve clarity and maintainability across the agent's functionality. - Enhanced error handling and cache management in the agent's initialization process. --- agent/{ => api}/agent.go | 15 +- agent/{ => api}/api.go | 59 +-- agent/{ => api}/api_test.go | 2 +- agent/assistant/api.go | 2 +- agent/context/chat.go | 182 +++++++++ agent/context/chat_test.go | 375 ++++++++++++++++++ agent/context/context.go | 6 - agent/context/context_test.go | 21 +- agent/context/gin.go | 81 ---- agent/context/message.go | 96 +++++ agent/context/message_test.go | 387 +++++++++++++++++++ agent/context/openapi.go | 313 +++++++++++++++ agent/context/openapi_chat_test.go | 586 +++++++++++++++++++++++++++++ agent/context/types.go | 97 ++++- agent/context/utils.go | 1 - agent/load.go | 55 +-- agent/process.go | 14 +- agent/types/dsl.go | 11 + agent/{ => types}/types.go | 6 +- openapi/agent/agent.go | 17 +- openapi/chat/completions.go | 41 +- openapi/oauth/guard.go | 8 +- service/service.go | 2 +- test/utils.go | 1 + widgets/app/app.go | 19 +- 25 files changed, 2197 insertions(+), 200 deletions(-) rename agent/{ => api}/agent.go (64%) rename agent/{ => api}/api.go (94%) rename agent/{ => api}/api_test.go (99%) create mode 100644 agent/context/chat.go create mode 100644 agent/context/chat_test.go delete mode 100644 agent/context/gin.go create mode 100644 agent/context/message.go create mode 100644 agent/context/message_test.go create mode 100644 agent/context/openapi.go create mode 100644 agent/context/openapi_chat_test.go create mode 100644 agent/types/dsl.go rename agent/{ => types}/types.go (95%) diff --git a/agent/agent.go b/agent/api/agent.go similarity index 64% rename from agent/agent.go rename to agent/api/agent.go index a89764b3..d41e8d80 100644 --- a/agent/agent.go +++ b/agent/api/agent.go @@ -1,13 +1,22 @@ -package agent +package api import ( "github.com/gin-gonic/gin" "github.com/yaoapp/yao/agent/assistant" chatctx "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/types" ) +// Agent the agent AI assistant +var Agent *API + +// API the agent API +type API struct { + *types.DSL +} + // Answer reply the message -func (agent *DSL) Answer(ctx chatctx.Context, question string, c *gin.Context) error { +func (agent *API) Answer(ctx chatctx.Context, question string, c *gin.Context) error { var err error var ast assistant.API = Agent.Assistant if ctx.AssistantID != "" { @@ -21,7 +30,7 @@ func (agent *DSL) Answer(ctx chatctx.Context, question string, c *gin.Context) e } // Select select an assistant -func (agent *DSL) Select(id string) (assistant.API, error) { +func (agent *API) Select(id string) (assistant.API, error) { if id == "" { return Agent.Assistant, nil } diff --git a/agent/api.go b/agent/api/api.go similarity index 94% rename from agent/api.go rename to agent/api/api.go index 1962608a..3fe8ac21 100644 --- a/agent/api.go +++ b/agent/api/api.go @@ -1,4 +1,4 @@ -package agent +package api import ( "fmt" @@ -14,12 +14,13 @@ import ( chatctx "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/message" store "github.com/yaoapp/yao/agent/store/types" + "github.com/yaoapp/yao/agent/types" "github.com/yaoapp/yao/helper" "github.com/yaoapp/yao/openapi/oauth" ) // API registers the Agent API endpoints -func (agent *DSL) API(router *gin.Engine, path string) error { +func (agent *API) API(router *gin.Engine, path string) error { // Get the guards middlewares, err := agent.getGuardHandlers() @@ -145,13 +146,13 @@ func (agent *DSL) API(router *gin.Engine, path string) error { } // handleStatus handles the status request -func (agent *DSL) handleStatus(c *gin.Context) { +func (agent *API) handleStatus(c *gin.Context) { c.Status(200) c.Done() } // handleChat handles the chat request -func (agent *DSL) handleChat(c *gin.Context) { +func (agent *API) handleChat(c *gin.Context) { // Set headers for SSE c.Header("Content-Type", "text/event-stream;charset=utf-8") c.Header("Cache-Control", "no-cache") @@ -215,7 +216,7 @@ func (agent *DSL) handleChat(c *gin.Context) { } // handleChatList handles the chat list request -func (agent *DSL) handleChatList(c *gin.Context) { +func (agent *API) handleChatList(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -259,7 +260,7 @@ func (agent *DSL) handleChatList(c *gin.Context) { } // handleChatHistory handles the chat history request -func (agent *DSL) handleChatHistory(c *gin.Context) { +func (agent *API) handleChatHistory(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -280,7 +281,7 @@ func (agent *DSL) handleChatHistory(c *gin.Context) { } // getOrigin returns the request origin -func (agent *DSL) getOrigin(c *gin.Context) string { +func (agent *API) getOrigin(c *gin.Context) string { origin := c.Request.Header.Get("Origin") if origin == "" { origin = c.Request.Referer() @@ -294,12 +295,12 @@ func (agent *DSL) getOrigin(c *gin.Context) string { } // getGuardHandlers returns authentication middleware handlers -func (agent *DSL) getGuardHandlers() ([]gin.HandlerFunc, error) { +func (agent *API) getGuardHandlers() ([]gin.HandlerFunc, error) { return []gin.HandlerFunc{}, nil } // defaultGuard is the default authentication handler -func (agent *DSL) defaultGuard(c *gin.Context) { +func (agent *API) defaultGuard(c *gin.Context) { // Check if the request is for OpenAPI OAuth if oauth.OAuth != nil { @@ -320,7 +321,7 @@ func (agent *DSL) defaultGuard(c *gin.Context) { } // Openapi Oauth -func (agent *DSL) guardOpenapiOauth(c *gin.Context) { +func (agent *API) guardOpenapiOauth(c *gin.Context) { s := oauth.OAuth token := agent.getAccessToken(c) if token == "" { @@ -348,7 +349,7 @@ func (agent *DSL) guardOpenapiOauth(c *gin.Context) { c.Set("__sid", sid) } -func (agent *DSL) getAccessToken(c *gin.Context) string { +func (agent *API) getAccessToken(c *gin.Context) string { token := c.GetHeader("Authorization") if token == "" || token == "Bearer undefined" { cookie, err := c.Cookie("__Host-access_token") @@ -360,7 +361,7 @@ func (agent *DSL) getAccessToken(c *gin.Context) string { return strings.TrimPrefix(token, "Bearer ") } -func (agent *DSL) getSessionID(c *gin.Context) string { +func (agent *API) getSessionID(c *gin.Context) string { sid, err := c.Cookie("__Host-session_id") if err != nil { return "" @@ -369,7 +370,7 @@ func (agent *DSL) getSessionID(c *gin.Context) string { } // handleChatLatest handles getting the latest chat -func (agent *DSL) handleChatLatest(c *gin.Context) { +func (agent *API) handleChatLatest(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -454,7 +455,7 @@ func (agent *DSL) handleChatLatest(c *gin.Context) { } // handleChatDetail handles getting a single chat's details -func (agent *DSL) handleChatDetail(c *gin.Context) { +func (agent *API) handleChatDetail(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -503,7 +504,7 @@ func (agent *DSL) handleChatDetail(c *gin.Context) { } // handleMentions handles getting mentions for a chat -func (agent *DSL) handleMentions(c *gin.Context) { +func (agent *API) handleMentions(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -536,9 +537,9 @@ func (agent *DSL) handleMentions(c *gin.Context) { } // Convert assistants to mentions - mentions := []Mention{} + mentions := []types.Mention{} for _, assistant := range response.Data { - mention := Mention{ + mention := types.Mention{ ID: assistant.ID, Name: assistant.Name, Type: assistant.Type, @@ -552,7 +553,7 @@ func (agent *DSL) handleMentions(c *gin.Context) { } // handleChatUpdate handles updating a chat's details -func (agent *DSL) handleChatUpdate(c *gin.Context) { +func (agent *API) handleChatUpdate(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -596,7 +597,7 @@ func (agent *DSL) handleChatUpdate(c *gin.Context) { } // handleChatDelete handles deleting a single chat -func (agent *DSL) handleChatDelete(c *gin.Context) { +func (agent *API) handleChatDelete(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -623,7 +624,7 @@ func (agent *DSL) handleChatDelete(c *gin.Context) { } // handleChatsDeleteAll handles deleting all chats for a user -func (agent *DSL) handleChatsDeleteAll(c *gin.Context) { +func (agent *API) handleChatsDeleteAll(c *gin.Context) { sid := c.GetString("__sid") if sid == "" { c.JSON(400, gin.H{"message": "sid is required", "code": 400}) @@ -643,7 +644,7 @@ func (agent *DSL) handleChatsDeleteAll(c *gin.Context) { } // handleGenerateTitle handles generating a chat title -func (agent *DSL) handleGenerateTitle(c *gin.Context) { +func (agent *API) handleGenerateTitle(c *gin.Context) { // Set headers for SSE c.Header("Content-Type", "text/event-stream;charset=utf-8") c.Header("Cache-Control", "no-cache") @@ -683,7 +684,7 @@ func (agent *DSL) handleGenerateTitle(c *gin.Context) { } // handleGeneratePrompts handles generating prompts -func (agent *DSL) handleGeneratePrompts(c *gin.Context) { +func (agent *API) handleGeneratePrompts(c *gin.Context) { // Set headers for SSE c.Header("Content-Type", "text/event-stream;charset=utf-8") c.Header("Cache-Control", "no-cache") @@ -722,7 +723,7 @@ func (agent *DSL) handleGeneratePrompts(c *gin.Context) { } // HandleAssistantList handles listing assistants (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantList(c *gin.Context) { +func (agent *API) HandleAssistantList(c *gin.Context) { // Parse filter parameters filter := store.AssistantFilter{ Type: "assistant", @@ -825,7 +826,7 @@ func parseBoolValue(value string) *bool { } // HandleAssistantCall handles the assistant API call (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantCall(c *gin.Context) { +func (agent *API) HandleAssistantCall(c *gin.Context) { assistantID := c.Param("id") if assistantID == "" { c.JSON(400, gin.H{"message": "assistant id is required", "code": 400}) @@ -865,7 +866,7 @@ func (agent *DSL) HandleAssistantCall(c *gin.Context) { } // HandleAssistantDetail handles getting a single assistant's details (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantDetail(c *gin.Context) { +func (agent *API) HandleAssistantDetail(c *gin.Context) { assistantID := c.Param("id") if assistantID == "" { c.JSON(400, gin.H{"message": "assistant id is required", "code": 400}) @@ -904,7 +905,7 @@ func (agent *DSL) HandleAssistantDetail(c *gin.Context) { } // HandleAssistantSave handles creating or updating an assistant (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantSave(c *gin.Context) { +func (agent *API) HandleAssistantSave(c *gin.Context) { var assistantData map[string]interface{} if err := c.BindJSON(&assistantData); err != nil { c.JSON(400, gin.H{"message": "invalid request body", "code": 400}) @@ -950,7 +951,7 @@ func (agent *DSL) HandleAssistantSave(c *gin.Context) { } // HandleAssistantDelete handles deleting an assistant (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantDelete(c *gin.Context) { +func (agent *API) HandleAssistantDelete(c *gin.Context) { assistantID := c.Param("id") if assistantID == "" { c.JSON(400, gin.H{"message": "assistant id is required", "code": 400}) @@ -976,7 +977,7 @@ func (agent *DSL) HandleAssistantDelete(c *gin.Context) { } // handleConnectors handles listing connectors -func (agent *DSL) handleConnectors(c *gin.Context) { +func (agent *API) handleConnectors(c *gin.Context) { options := []map[string]interface{}{} // Filter and format connectors @@ -1002,7 +1003,7 @@ func (agent *DSL) handleConnectors(c *gin.Context) { } // HandleAssistantTags handles getting all assistant tags (exported for use in openapi/agent) -func (agent *DSL) HandleAssistantTags(c *gin.Context) { +func (agent *API) HandleAssistantTags(c *gin.Context) { locale := "en-us" // Default locale if loc := c.Query("locale"); loc != "" { locale = strings.ToLower(strings.TrimSpace(loc)) diff --git a/agent/api_test.go b/agent/api/api_test.go similarity index 99% rename from agent/api_test.go rename to agent/api/api_test.go index e10b4a0f..20ede75f 100644 --- a/agent/api_test.go +++ b/agent/api/api_test.go @@ -1,4 +1,4 @@ -package agent +package api // import ( // "context" diff --git a/agent/assistant/api.go b/agent/assistant/api.go index 5f7f5882..f6f6f4bb 100644 --- a/agent/assistant/api.go +++ b/agent/assistant/api.go @@ -302,7 +302,7 @@ func (ast *Assistant) Call(c *gin.Context, payload APIPayload) (interface{}, err return nil, fmt.Errorf(HookErrorMethodNotFound) } - if payload.Args == nil || len(payload.Args) == 0 { + if len(payload.Args) == 0 { return scriptCtx.CallWith(ctx, method) } diff --git a/agent/context/chat.go b/agent/context/chat.go new file mode 100644 index 00000000..2e4f22be --- /dev/null +++ b/agent/context/chat.go @@ -0,0 +1,182 @@ +package context + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "time" + + gonanoid "github.com/matoous/go-nanoid/v2" + "github.com/yaoapp/gou/store" +) + +const ( + chatCachePrefix = "chat:messages:" + chatCacheTTL = time.Hour * 24 * 7 // 7 days +) + +// filterNonAssistantMessages returns messages excluding assistant messages +func filterNonAssistantMessages(messages []Message) []Message { + var filtered []Message + for _, msg := range messages { + if msg.Role != RoleAssistant { + filtered = append(filtered, msg) + } + } + return filtered +} + +// countUserMessages returns the number of user role messages +func countUserMessages(messages []Message) int { + count := 0 + for _, msg := range messages { + if msg.Role == RoleUser { + count++ + } + } + return count +} + +// GetChatIDByMessages gets or generates a chat ID based on message content +// Matching strategy: +// - Only non-assistant messages (system, developer, user, tool) are used for matching +// - User adds a new message at the end each time +// - To detect continuation, we match messages BEFORE the last non-assistant message +// - If previous non-assistant messages match cached conversation → same chat +// - For first user message (even with system/developer messages): always generate new chat ID +func GetChatIDByMessages(cache store.Store, messages []Message) (string, error) { + if len(messages) == 0 { + return "", fmt.Errorf("messages cannot be empty") + } + + // Filter out assistant messages for matching + nonAssistantMessages := filterNonAssistantMessages(messages) + + // Count user messages to determine matching strategy + userMessageCount := countUserMessages(nonAssistantMessages) + + var chatID string + var matched bool + + // Matching strategy based on user message count: + // - 1 user message: generate new chat ID (cannot determine continuation) + // - 2+ user messages: match all except last (which is the new user input) + if userMessageCount >= 2 { + // Match previous messages (all except last non-assistant message) + matchMessages := nonAssistantMessages[:len(nonAssistantMessages)-1] + hash, err := hashMessages(matchMessages) + if err == nil { + key := getKey(hash) + if cachedID, ok := cache.Get(key); ok { + if chatIDStr, ok := cachedID.(string); ok && chatIDStr != "" { + chatID = chatIDStr + matched = true + } + } + } + } + + // If no match, generate new chat ID + if !matched { + chatID = GenChatID() + } + + // Cache the current messages for future matching + // CacheChatID will handle filtering assistant messages + // Next request will have one more message and will try to match current messages + _ = CacheChatID(cache, messages, chatID) + + return chatID, nil +} + +// CacheChatID cache the chat ID with all message prefixes for future matching +// It caches ALL prefixes of the message array to enable conversation continuation detection +// Assistant messages are automatically filtered out before caching +// Example: For messages [A,B,C], it caches hashes for [A], [A,B], and [A,B,C] +func CacheChatID(cache store.Store, messages []Message, chatID string) error { + if len(messages) == 0 { + return fmt.Errorf("messages cannot be empty") + } + + if chatID == "" { + return fmt.Errorf("chatID cannot be empty") + } + + // Filter out assistant messages + nonAssistantMessages := filterNonAssistantMessages(messages) + if len(nonAssistantMessages) == 0 { + return fmt.Errorf("no non-assistant messages to cache") + } + + // Cache all prefixes of the non-assistant messages array + // This allows detecting conversation continuation when new messages are added + for length := 1; length <= len(nonAssistantMessages); length++ { + prefix := nonAssistantMessages[:length] + hash, err := hashMessages(prefix) + if err != nil { + continue // Skip this prefix if hashing fails + } + + key := getKey(hash) + // Ignore errors for individual cache sets + _ = cache.Set(key, chatID, chatCacheTTL) + } + + return nil +} + +// GenChatID generate a new chat ID using NanoID algorithm +// safe: optional parameter, reserved for future safe mode implementation (collision detection) +func GenChatID(safe ...bool) string { + // TODO: Implement safe mode with collision detection when needed + // For now, NanoID provides sufficient uniqueness without collision checking + + // URL-safe alphabet (no ambiguous characters like 0/O, 1/l/I) + const alphabet = "23456789ABCDEFGHJKMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz" + const length = 16 // 16 characters provides good balance of uniqueness and readability + + id, err := gonanoid.Generate(alphabet, length) + if err != nil { + // Fallback to timestamp-based ID if NanoID generation fails + return fmt.Sprintf("%d", time.Now().UnixNano()) + } + + return id +} + +// getKey generates a cache key for messages +func getKey(messageHash string) string { + return chatCachePrefix + messageHash +} + +// hashMessage generates a hash for a single message +func hashMessage(msg Message) (string, error) { + data, err := json.Marshal(msg) + if err != nil { + return "", err + } + hash := sha256.Sum256(data) + return hex.EncodeToString(hash[:]), nil +} + +// hashMessages generates a combined hash for a slice of messages +// Note: Caller is responsible for filtering messages (e.g., removing assistant messages) +func hashMessages(messages []Message) (string, error) { + if len(messages) == 0 { + return "", fmt.Errorf("messages cannot be empty") + } + + var hashes string + for _, msg := range messages { + hash, err := hashMessage(msg) + if err != nil { + return "", err + } + hashes += hash + } + + // Generate final hash from combined hashes + finalHash := sha256.Sum256([]byte(hashes)) + return hex.EncodeToString(finalHash[:]), nil +} diff --git a/agent/context/chat_test.go b/agent/context/chat_test.go new file mode 100644 index 00000000..f9adbf9f --- /dev/null +++ b/agent/context/chat_test.go @@ -0,0 +1,375 @@ +package context + +import ( + "testing" + + "github.com/yaoapp/gou/store" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/test" +) + +func getTestCache(t *testing.T) store.Store { + cache, err := store.Get("__yao.agent.cache") + if err != nil { + t.Fatalf("Failed to get cache store: %v", err) + } + cache.Clear() // Clean before test + return cache +} + +func TestGetChatIDByMessages_NewConversation(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + messages := []Message{ + { + Role: RoleUser, + Content: "Hello, how are you?", + }, + } + + // First request - should generate new chat ID + chatID1, err := GetChatIDByMessages(cache, messages) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID1 == "" { + t.Fatal("Expected non-empty chat ID") + } + + // Second request with same single user message - should generate DIFFERENT chat ID + // (single user message always generates new chat ID to avoid false matches) + chatID2, err := GetChatIDByMessages(cache, messages) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID2 == "" { + t.Fatal("Expected non-empty chat ID") + } + + // Both should be valid but different (single user message = new conversation each time) + if chatID1 == chatID2 { + t.Errorf("Expected different chat IDs for single user message, got same ID: %s", chatID1) + } +} + +func TestGetChatIDByMessages_ContinuousConversation(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + // Scenario: User conversation with incrementally added messages + // Request 1: [user1] + messages1 := []Message{ + {Role: RoleUser, Content: "First message"}, + } + chatID1, err := GetChatIDByMessages(cache, messages1) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + // Request 2: [user1, user2] + // For 2 messages, matches last 1 message + // Should match chatID1 because last message is cached + messages2 := []Message{ + {Role: RoleUser, Content: "First message"}, + {Role: RoleUser, Content: "Second message"}, + } + chatID2, err := GetChatIDByMessages(cache, messages2) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID1 != chatID2 { + t.Errorf("Expected chatID2 to match chatID1, got %s and %s", chatID2, chatID1) + } + + // Request 3: [user1, user2, user3] + // For 3+ messages, matches last 2 messages + // Should match chatID2 because last 2 messages are cached + messages3 := []Message{ + {Role: RoleUser, Content: "First message"}, + {Role: RoleUser, Content: "Second message"}, + {Role: RoleUser, Content: "Third message"}, + } + chatID3, err := GetChatIDByMessages(cache, messages3) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID2 != chatID3 { + t.Errorf("Expected chatID3 to match chatID2, got %s and %s", chatID3, chatID2) + } + + // Request 4: [user1, user2, user3, user4] + // Should match chatID3 because last 2 messages are cached + messages4 := []Message{ + {Role: RoleUser, Content: "First message"}, + {Role: RoleUser, Content: "Second message"}, + {Role: RoleUser, Content: "Third message"}, + {Role: RoleUser, Content: "Fourth message"}, + } + chatID4, err := GetChatIDByMessages(cache, messages4) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID3 != chatID4 { + t.Errorf("Expected chatID4 to match chatID3, got %s and %s", chatID4, chatID3) + } + + // All should be the same conversation + if chatID1 != chatID4 { + t.Errorf("Expected all chat IDs to be the same, got %s and %s", chatID1, chatID4) + } +} + +func TestGetChatIDByMessages_DifferentConversations(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + // First conversation + messages1 := []Message{ + { + Role: RoleUser, + Content: "Hello", + }, + } + + chatID1, err := GetChatIDByMessages(cache, messages1) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + err = CacheChatID(cache, messages1, chatID1) + if err != nil { + t.Fatalf("Failed to cache chat ID: %v", err) + } + + // Different conversation + messages2 := []Message{ + { + Role: RoleUser, + Content: "Goodbye", + }, + } + + chatID2, err := GetChatIDByMessages(cache, messages2) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID1 == chatID2 { + t.Errorf("Expected different chat IDs for different conversations, got %s", chatID1) + } +} + +func TestGetChatIDByMessages_MultiModalContent(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + // First request with multimodal content + messages1 := []Message{ + { + Role: RoleUser, + Content: []ContentPart{ + { + Type: ContentText, + Text: "What's in this image?", + }, + { + Type: ContentImageURL, + ImageURL: &ImageURL{ + URL: "https://example.com/image.jpg", + Detail: DetailHigh, + }, + }, + }, + }, + } + + chatID1, err := GetChatIDByMessages(cache, messages1) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + // Second request - add another message to continue conversation + messages2 := append(messages1, Message{ + Role: RoleUser, + Content: "Tell me more details", + }) + + chatID2, err := GetChatIDByMessages(cache, messages2) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + // Should get same chat ID (continuation) + if chatID1 != chatID2 { + t.Errorf("Expected same chat ID for multimodal continuation, got %s and %s", chatID1, chatID2) + } +} + +func TestGetChatIDByMessages_WithToolCalls(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + // First request with user message + messages1 := []Message{ + { + Role: RoleUser, + Content: "What's the weather in Tokyo?", + }, + } + + chatID1, err := GetChatIDByMessages(cache, messages1) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + // Second request - add assistant response and another user message + messages2 := []Message{ + { + Role: RoleUser, + Content: "What's the weather in Tokyo?", + }, + { + Role: RoleAssistant, + Content: nil, + ToolCalls: []ToolCall{ + { + ID: "call_123", + Type: ToolTypeFunction, + Function: Function{ + Name: "get_weather", + Arguments: `{"location":"Tokyo"}`, + }, + }, + }, + }, + { + Role: RoleUser, + Content: "How about tomorrow?", + }, + } + + chatID2, err := GetChatIDByMessages(cache, messages2) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + // Should get same chat ID (assistant messages are ignored, so it matches the first user message) + if chatID1 != chatID2 { + t.Errorf("Expected same chat ID for messages with tool calls, got %s and %s", chatID1, chatID2) + } +} + +func TestCacheChatID_EmptyMessages(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + err := CacheChatID(cache, []Message{}, "chat_123") + if err == nil { + t.Error("Expected error for empty messages") + } +} + +func TestCacheChatID_EmptyChatID(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + messages := []Message{ + { + Role: RoleUser, + Content: "Hello", + }, + } + + err := CacheChatID(cache, messages, "") + if err == nil { + t.Error("Expected error for empty chat ID") + } +} + +func TestGetChatIDByMessages_EmptyMessages(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + cache := getTestCache(t) + + _, err := GetChatIDByMessages(cache, []Message{}) + if err == nil { + t.Error("Expected error for empty messages") + } +} + +func TestHashMessage_Consistency(t *testing.T) { + msg := Message{ + Role: RoleUser, + Content: "Test message", + } + + hash1, err := hashMessage(msg) + if err != nil { + t.Fatalf("Failed to hash message: %v", err) + } + + hash2, err := hashMessage(msg) + if err != nil { + t.Fatalf("Failed to hash message: %v", err) + } + + if hash1 != hash2 { + t.Errorf("Expected consistent hashes, got %s and %s", hash1, hash2) + } +} + +func TestGetKey(t *testing.T) { + hash := "abc123" + key := getKey(hash) + + expectedPrefix := chatCachePrefix + if len(key) <= len(expectedPrefix) { + t.Errorf("Expected key to have prefix, got %s", key) + } + + if key[:len(expectedPrefix)] != expectedPrefix { + t.Errorf("Expected key to start with %s, got %s", expectedPrefix, key) + } + + if key != chatCachePrefix+hash { + t.Errorf("Expected key %s, got %s", chatCachePrefix+hash, key) + } +} + +func TestGenChatID(t *testing.T) { + id1 := GenChatID() + + if id1 == "" { + t.Error("Expected non-empty chat ID") + } + + // Check length - NanoID with length 16 should produce 16 character strings + if len(id1) < 10 { + t.Errorf("Expected chat ID to have reasonable length, got %d characters: %s", len(id1), id1) + } + + // Note: We don't test uniqueness here because nano timestamp-based IDs + // can occasionally be the same when generated in rapid succession. + // The uniqueness is good enough for production use. +} diff --git a/agent/context/context.go b/agent/context/context.go index e9cb15bc..ed418b3b 100644 --- a/agent/context/context.go +++ b/agent/context/context.go @@ -2,7 +2,6 @@ package context import ( "context" - "fmt" "time" jsoniter "github.com/json-iterator/go" @@ -126,8 +125,3 @@ func (ctx *Context) Map() map[string]interface{} { return data } - -// GenChatID generate a new chat ID -func GenChatID() string { - return fmt.Sprintf("chat_%d", time.Now().UnixNano()) -} diff --git a/agent/context/context_test.go b/agent/context/context_test.go index 27200ef0..92de88d0 100644 --- a/agent/context/context_test.go +++ b/agent/context/context_test.go @@ -9,7 +9,7 @@ import ( "github.com/stretchr/testify/assert" ) -func TestNewGin(t *testing.T) { +func TestNewOpenAPI(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { @@ -28,21 +28,22 @@ func TestNewGin(t *testing.T) { { name: "Parse all query parameters", queryParams: map[string]string{ - "chat_id": "chat123", - "locale": "zh-CN", - "theme": "dark", - "referer": RefererProcess, - "accept": string(AcceptStandard), + "assistant_id": "ast456", + "chat_id": "chat123", + "locale": "zh-CN", + "theme": "Dark", + "referer": RefererProcess, + "accept": string(AcceptStandard), }, routeParams: map[string]string{ - "assistant_id": "ast456", + "assistant_id": "route-ast", }, headers: map[string]string{ "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)", }, expectedChatID: "chat123", expectedAssistant: "ast456", - expectedLocale: "zh-CN", + expectedLocale: "zh-cn", expectedTheme: "dark", expectedClientType: "macos", expectedReferer: RefererProcess, @@ -168,7 +169,7 @@ func TestNewGin(t *testing.T) { } // Call NewGin - ctx := NewGin(c) + ctx := NewOpenAPI(c, nil) // Assertions assert.Equal(t, tt.expectedChatID, ctx.ChatID, "ChatID mismatch") @@ -204,7 +205,7 @@ func TestParseClientType(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - result := parseClientType(tt.userAgent) + result := getClientType(tt.userAgent) assert.Equal(t, tt.expected, result) }) } diff --git a/agent/context/gin.go b/agent/context/gin.go deleted file mode 100644 index 8159e3a9..00000000 --- a/agent/context/gin.go +++ /dev/null @@ -1,81 +0,0 @@ -package context - -import ( - "strings" - - "github.com/gin-gonic/gin" - "github.com/yaoapp/gou/plan" - "github.com/yaoapp/yao/openapi/oauth/authorized" -) - -// NewGin create a new context from gin context -func NewGin(c *gin.Context) Context { - // Get authorized information - authInfo := authorized.GetInfo(c) - - // Extract parameters from query and route - chatID := c.Query("chat_id") - assistantID := c.Param("assistant_id") // Get from route parameter - locale := c.Query("locale") - theme := c.Query("theme") - referer := c.Query("referer") - accept := c.Query("accept") - - // Parse client information from User-Agent header - userAgent := c.GetHeader("User-Agent") - clientType := parseClientType(userAgent) - clientIP := c.ClientIP() - - // Create base context - ctx := Context{ - Context: c.Request.Context(), - Space: plan.NewMemorySharedSpace(), - Authorized: authInfo, - ChatID: chatID, - AssistantID: assistantID, - Locale: locale, - Theme: theme, - Client: Client{ - Type: clientType, - UserAgent: userAgent, - IP: clientIP, - }, - } - - // Get Referer from query parameter, header, or default - ctx.Referer = getValidatedValue(referer, c.GetHeader("X-Yao-Referer"), RefererAPI, validateReferer) - - // Get Accept from query parameter, header, or default - ctx.Accept = getValidatedAccept(accept, c.GetHeader("X-Yao-Accept"), clientType) - - return ctx -} - -// parseClientType parses the client type from User-Agent header -func parseClientType(userAgent string) string { - if userAgent == "" { - return "web" // Default to web - } - - ua := strings.ToLower(userAgent) - - // Check for specific client types - switch { - case strings.Contains(ua, "yao-agent") || strings.Contains(ua, "agent"): - return "agent" - case strings.Contains(ua, "yao-jssdk") || strings.Contains(ua, "jssdk"): - return "jssdk" - case strings.Contains(ua, "android"): - return "android" - case strings.Contains(ua, "iphone") || strings.Contains(ua, "ipad") || strings.Contains(ua, "ipod"): - return "ios" - case strings.Contains(ua, "windows"): - return "windows" - case strings.Contains(ua, "mac os x") || strings.Contains(ua, "macintosh"): - return "macos" - case strings.Contains(ua, "linux"): - return "linux" - default: - return "web" - } -} diff --git a/agent/context/message.go b/agent/context/message.go new file mode 100644 index 00000000..b3c716c8 --- /dev/null +++ b/agent/context/message.go @@ -0,0 +1,96 @@ +package context + +import ( + "encoding/json" + "fmt" +) + +// UnmarshalJSON custom unmarshaler for Message to handle Content field +func (m *Message) UnmarshalJSON(data []byte) error { + // Define a temporary struct to avoid infinite recursion + type Alias Message + aux := &struct { + Content json.RawMessage `json:"content,omitempty"` + *Alias + }{ + Alias: (*Alias)(m), + } + + if err := json.Unmarshal(data, &aux); err != nil { + return err + } + + // If content is empty, return early + if len(aux.Content) == 0 || string(aux.Content) == "null" { + m.Content = nil + return nil + } + + // Try to unmarshal as string first + var contentStr string + if err := json.Unmarshal(aux.Content, &contentStr); err == nil { + m.Content = contentStr + return nil + } + + // Try to unmarshal as array of ContentPart + var contentParts []ContentPart + if err := json.Unmarshal(aux.Content, &contentParts); err == nil { + m.Content = contentParts + return nil + } + + return fmt.Errorf("content must be either a string or an array of ContentPart") +} + +// MarshalJSON custom marshaler for Message +func (m *Message) MarshalJSON() ([]byte, error) { + type Alias Message + return json.Marshal(&struct { + *Alias + }{ + Alias: (*Alias)(m), + }) +} + +// NewTextMessage creates a new message with text content +func NewTextMessage(role MessageRole, text string) *Message { + return &Message{ + Role: role, + Content: text, + } +} + +// NewMultipartMessage creates a new message with multipart content +func NewMultipartMessage(role MessageRole, parts []ContentPart) *Message { + return &Message{ + Role: role, + Content: parts, + } +} + +// GetContentAsString returns content as string if possible +func (m *Message) GetContentAsString() (string, bool) { + if str, ok := m.Content.(string); ok { + return str, true + } + return "", false +} + +// GetContentAsParts returns content as ContentPart array if possible +func (m *Message) GetContentAsParts() ([]ContentPart, bool) { + if parts, ok := m.Content.([]ContentPart); ok { + return parts, true + } + return nil, false +} + +// HasToolCalls checks if the message has tool calls +func (m *Message) HasToolCalls() bool { + return len(m.ToolCalls) > 0 +} + +// IsRefusal checks if the message is a refusal +func (m *Message) IsRefusal() bool { + return m.Refusal != nil && *m.Refusal != "" +} diff --git a/agent/context/message_test.go b/agent/context/message_test.go new file mode 100644 index 00000000..7c20b712 --- /dev/null +++ b/agent/context/message_test.go @@ -0,0 +1,387 @@ +package context + +import ( + "encoding/json" + "testing" +) + +func TestMessage_UnmarshalJSON_StringContent(t *testing.T) { + jsonData := `{ + "role": "user", + "content": "Hello, world!" + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + if msg.Role != RoleUser { + t.Errorf("Expected role %s, got %s", RoleUser, msg.Role) + } + + content, ok := msg.GetContentAsString() + if !ok { + t.Fatal("Expected content to be string") + } + + if content != "Hello, world!" { + t.Errorf("Expected content 'Hello, world!', got '%s'", content) + } +} + +func TestMessage_UnmarshalJSON_ArrayContent(t *testing.T) { + jsonData := `{ + "role": "user", + "content": [ + { + "type": "text", + "text": "What's in this image?" + }, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/image.jpg", + "detail": "high" + } + } + ] + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + if msg.Role != RoleUser { + t.Errorf("Expected role %s, got %s", RoleUser, msg.Role) + } + + parts, ok := msg.GetContentAsParts() + if !ok { + t.Fatal("Expected content to be array of ContentPart") + } + + if len(parts) != 2 { + t.Fatalf("Expected 2 content parts, got %d", len(parts)) + } + + // Check first part (text) + if parts[0].Type != ContentText { + t.Errorf("Expected type %s, got %s", ContentText, parts[0].Type) + } + if parts[0].Text != "What's in this image?" { + t.Errorf("Expected text 'What's in this image?', got '%s'", parts[0].Text) + } + + // Check second part (image) + if parts[1].Type != ContentImageURL { + t.Errorf("Expected type %s, got %s", ContentImageURL, parts[1].Type) + } + if parts[1].ImageURL == nil { + t.Fatal("Expected ImageURL to be non-nil") + } + if parts[1].ImageURL.URL != "https://example.com/image.jpg" { + t.Errorf("Expected URL 'https://example.com/image.jpg', got '%s'", parts[1].ImageURL.URL) + } + if parts[1].ImageURL.Detail != DetailHigh { + t.Errorf("Expected detail %s, got %s", DetailHigh, parts[1].ImageURL.Detail) + } +} + +func TestMessage_UnmarshalJSON_NullContent(t *testing.T) { + jsonData := `{ + "role": "assistant", + "content": null, + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": "{\"location\":\"Tokyo\"}" + } + } + ] + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + if msg.Role != RoleAssistant { + t.Errorf("Expected role %s, got %s", RoleAssistant, msg.Role) + } + + if msg.Content != nil { + t.Errorf("Expected content to be nil, got %v", msg.Content) + } + + if !msg.HasToolCalls() { + t.Fatal("Expected message to have tool calls") + } + + if len(msg.ToolCalls) != 1 { + t.Fatalf("Expected 1 tool call, got %d", len(msg.ToolCalls)) + } + + if msg.ToolCalls[0].ID != "call_123" { + t.Errorf("Expected tool call ID 'call_123', got '%s'", msg.ToolCalls[0].ID) + } +} + +func TestMessage_UnmarshalJSON_WithRefusal(t *testing.T) { + refusalText := "I cannot help with that request." + jsonData := `{ + "role": "assistant", + "content": "I'm sorry, but I can't assist with that.", + "refusal": "I cannot help with that request." + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + if !msg.IsRefusal() { + t.Error("Expected message to be a refusal") + } + + if msg.Refusal == nil { + t.Fatal("Expected refusal to be non-nil") + } + + if *msg.Refusal != refusalText { + t.Errorf("Expected refusal '%s', got '%s'", refusalText, *msg.Refusal) + } +} + +func TestMessage_UnmarshalJSON_AudioContent(t *testing.T) { + jsonData := `{ + "role": "user", + "content": [ + { + "type": "text", + "text": "Transcribe this audio" + }, + { + "type": "input_audio", + "input_audio": { + "data": "base64encodedaudiodata", + "format": "wav" + } + } + ] + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + parts, ok := msg.GetContentAsParts() + if !ok { + t.Fatal("Expected content to be array of ContentPart") + } + + if len(parts) != 2 { + t.Fatalf("Expected 2 content parts, got %d", len(parts)) + } + + // Check audio part + if parts[1].Type != ContentInputAudio { + t.Errorf("Expected type %s, got %s", ContentInputAudio, parts[1].Type) + } + if parts[1].InputAudio == nil { + t.Fatal("Expected InputAudio to be non-nil") + } + if parts[1].InputAudio.Data != "base64encodedaudiodata" { + t.Errorf("Expected audio data 'base64encodedaudiodata', got '%s'", parts[1].InputAudio.Data) + } + if parts[1].InputAudio.Format != "wav" { + t.Errorf("Expected format 'wav', got '%s'", parts[1].InputAudio.Format) + } +} + +func TestMessage_MarshalJSON_StringContent(t *testing.T) { + msg := NewTextMessage(RoleUser, "Hello, AI!") + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("Failed to marshal: %v", err) + } + + var result map[string]interface{} + err = json.Unmarshal(data, &result) + if err != nil { + t.Fatalf("Failed to unmarshal result: %v", err) + } + + if result["role"] != string(RoleUser) { + t.Errorf("Expected role %s, got %v", RoleUser, result["role"]) + } + + if result["content"] != "Hello, AI!" { + t.Errorf("Expected content 'Hello, AI!', got %v", result["content"]) + } +} + +func TestMessage_MarshalJSON_ArrayContent(t *testing.T) { + parts := []ContentPart{ + { + Type: ContentText, + Text: "Describe this image", + }, + { + Type: ContentImageURL, + ImageURL: &ImageURL{ + URL: "https://example.com/test.jpg", + Detail: DetailLow, + }, + }, + } + + msg := NewMultipartMessage(RoleUser, parts) + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("Failed to marshal: %v", err) + } + + // Unmarshal back to verify + var result Message + err = json.Unmarshal(data, &result) + if err != nil { + t.Fatalf("Failed to unmarshal result: %v", err) + } + + resultParts, ok := result.GetContentAsParts() + if !ok { + t.Fatal("Expected content to be array of ContentPart") + } + + if len(resultParts) != 2 { + t.Fatalf("Expected 2 content parts, got %d", len(resultParts)) + } +} + +func TestMessage_MarshalJSON_WithToolCalls(t *testing.T) { + msg := &Message{ + Role: RoleAssistant, + Content: nil, + ToolCalls: []ToolCall{ + { + ID: "call_abc123", + Type: ToolTypeFunction, + Function: Function{ + Name: "get_weather", + Arguments: `{"location":"San Francisco"}`, + }, + }, + }, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("Failed to marshal: %v", err) + } + + // Unmarshal back to verify + var result Message + err = json.Unmarshal(data, &result) + if err != nil { + t.Fatalf("Failed to unmarshal result: %v", err) + } + + if !result.HasToolCalls() { + t.Error("Expected message to have tool calls") + } + + if len(result.ToolCalls) != 1 { + t.Fatalf("Expected 1 tool call, got %d", len(result.ToolCalls)) + } + + if result.ToolCalls[0].Function.Name != "get_weather" { + t.Errorf("Expected function name 'get_weather', got '%s'", result.ToolCalls[0].Function.Name) + } +} + +func TestMessage_ToolMessage(t *testing.T) { + toolCallID := "call_abc123" + jsonData := `{ + "role": "tool", + "tool_call_id": "call_abc123", + "content": "The weather in San Francisco is sunny, 72°F" + }` + + var msg Message + err := json.Unmarshal([]byte(jsonData), &msg) + if err != nil { + t.Fatalf("Failed to unmarshal: %v", err) + } + + if msg.Role != RoleTool { + t.Errorf("Expected role %s, got %s", RoleTool, msg.Role) + } + + if msg.ToolCallID == nil { + t.Fatal("Expected tool_call_id to be non-nil") + } + + if *msg.ToolCallID != toolCallID { + t.Errorf("Expected tool_call_id '%s', got '%s'", toolCallID, *msg.ToolCallID) + } + + content, ok := msg.GetContentAsString() + if !ok { + t.Fatal("Expected content to be string") + } + + if content != "The weather in San Francisco is sunny, 72°F" { + t.Errorf("Unexpected content: %s", content) + } +} + +func TestNewTextMessage(t *testing.T) { + msg := NewTextMessage(RoleSystem, "You are a helpful assistant.") + + if msg.Role != RoleSystem { + t.Errorf("Expected role %s, got %s", RoleSystem, msg.Role) + } + + content, ok := msg.GetContentAsString() + if !ok { + t.Fatal("Expected content to be string") + } + + if content != "You are a helpful assistant." { + t.Errorf("Expected content 'You are a helpful assistant.', got '%s'", content) + } +} + +func TestNewMultipartMessage(t *testing.T) { + parts := []ContentPart{ + {Type: ContentText, Text: "Hello"}, + } + + msg := NewMultipartMessage(RoleUser, parts) + + if msg.Role != RoleUser { + t.Errorf("Expected role %s, got %s", RoleUser, msg.Role) + } + + resultParts, ok := msg.GetContentAsParts() + if !ok { + t.Fatal("Expected content to be array of ContentPart") + } + + if len(resultParts) != 1 { + t.Fatalf("Expected 1 content part, got %d", len(resultParts)) + } +} diff --git a/agent/context/openapi.go b/agent/context/openapi.go new file mode 100644 index 00000000..b7af0b2c --- /dev/null +++ b/agent/context/openapi.go @@ -0,0 +1,313 @@ +package context + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "strings" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/gou/plan" + "github.com/yaoapp/gou/store" + "github.com/yaoapp/yao/openapi/oauth/authorized" +) + +// NewOpenAPI create a new context from openapi context +func NewOpenAPI(c *gin.Context, cache store.Store) Context { + // Get authorized information + authInfo := authorized.GetInfo(c) + + // Extract assistant ID (route parameter takes priority, handled in GetAssistantID) + assistantID, _ := GetAssistantID(c) + + // Extract chat ID (may generate from messages if not provided) + // GetChatID internally calls GetChatIDByMessages which auto-caches + chatID, _ := GetChatID(c, cache) + + // Parse client information from User-Agent header + userAgent := c.GetHeader("User-Agent") + clientType := getClientType(userAgent) + clientIP := c.ClientIP() + + // Create context with extracted parameters + ctx := Context{ + Context: c.Request.Context(), + Space: plan.NewMemorySharedSpace(), + Authorized: authInfo, + ChatID: chatID, + AssistantID: assistantID, + Locale: GetLocale(c), + Theme: GetTheme(c), + Referer: GetReferer(c), + Accept: GetAccept(c), + Client: Client{ + Type: clientType, + UserAgent: userAgent, + IP: clientIP, + }, + } + + return ctx +} + +// getClientType parses the client type from User-Agent header +func getClientType(userAgent string) string { + if userAgent == "" { + return "web" // Default to web + } + + ua := strings.ToLower(userAgent) + + // Check for specific client types + switch { + case strings.Contains(ua, "yao-agent") || strings.Contains(ua, "agent"): + return "agent" + case strings.Contains(ua, "yao-jssdk") || strings.Contains(ua, "jssdk"): + return "jssdk" + case strings.Contains(ua, "android"): + return "android" + case strings.Contains(ua, "iphone") || strings.Contains(ua, "ipad") || strings.Contains(ua, "ipod"): + return "ios" + case strings.Contains(ua, "windows"): + return "windows" + case strings.Contains(ua, "mac os x") || strings.Contains(ua, "macintosh"): + return "macos" + case strings.Contains(ua, "linux"): + return "linux" + default: + return "web" + } +} + +// getPayloadField reads a string field from request body +func getPayloadField(c *gin.Context, fieldName string) string { + if c.Request.Body == nil { + return "" + } + + body, err := io.ReadAll(c.Request.Body) + if err != nil { + return "" + } + + // Restore body for further use + c.Request.Body = io.NopCloser(bytes.NewReader(body)) + + if len(body) == 0 { + return "" + } + + // Parse JSON to extract field + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + return "" + } + + if value, ok := payload[fieldName]; ok { + if strValue, ok := value.(string); ok { + return strValue + } + } + + return "" +} + +// GetAssistantID extracts assistant ID from request with priority: +// 1. Query parameter "assistant_id" +// 2. Header "X-Yao-Assistant" +// 3. Query parameter "model" - splits by "-" takes last field, extracts ID from "yao_xxx" prefix +// 4. Payload "model" field - same parsing as query parameter +func GetAssistantID(c *gin.Context) (string, error) { + // Priority 1: Query parameter assistant_id + if assistantID := c.Query("assistant_id"); assistantID != "" { + return assistantID, nil + } + + // Priority 2: Header X-Yao-Assistant + if assistantID := c.GetHeader("X-Yao-Assistant"); assistantID != "" { + return assistantID, nil + } + + // Priority 3 & 4: Extract from model parameter (Query or Payload) + model := c.Query("model") + if model == "" { + model = getPayloadField(c, "model") + } + + if model != "" { + // Split by "-" and get the last field + parts := strings.Split(model, "-") + lastField := strings.TrimSpace(parts[len(parts)-1]) + + // Check if it has yao_ prefix + if strings.HasPrefix(lastField, "yao_") { + assistantID := strings.TrimPrefix(lastField, "yao_") + if assistantID != "" { + return assistantID, nil + } + } + } + + // If no assistant ID found, return error + return "", fmt.Errorf("assistant_id is required") +} + +// GetMessages extracts messages from the request +// Priority: +// 1. Query parameter "messages" (JSON string) +// 2. Request body "messages" field +func GetMessages(c *gin.Context) ([]Message, error) { + // Try query parameter first + if messagesJSON := c.Query("messages"); messagesJSON != "" { + var messages []Message + if err := json.Unmarshal([]byte(messagesJSON), &messages); err == nil && len(messages) > 0 { + return messages, nil + } + } + + // Check if request body exists + if c.Request.Body == nil { + return nil, fmt.Errorf("messages field is required") + } + + // Try request body + // Read body carefully to allow reuse + body, err := io.ReadAll(c.Request.Body) + if err != nil { + return nil, fmt.Errorf("failed to read request body: %w", err) + } + + // Restore body for further processing + c.Request.Body = io.NopCloser(bytes.NewBuffer(body)) + + // If body is empty, return error + if len(body) == 0 { + return nil, fmt.Errorf("messages field is required") + } + + var requestBody struct { + Messages []Message `json:"messages"` + } + + if err := json.Unmarshal(body, &requestBody); err != nil { + return nil, fmt.Errorf("failed to parse messages from request body: %w", err) + } + + if len(requestBody.Messages) == 0 { + return nil, fmt.Errorf("messages field is required and must not be empty") + } + + return requestBody.Messages, nil +} + +// GetLocale extracts locale from request with priority: +// 1. Query parameter "locale" +// 2. Header "Accept-Language" +func GetLocale(c *gin.Context) string { + // Priority 1: Query parameter + if locale := c.Query("locale"); locale != "" { + return strings.ToLower(locale) + } + + // Priority 2: Header Accept-Language + if acceptLang := c.GetHeader("Accept-Language"); acceptLang != "" { + // Parse Accept-Language header (e.g., "en-US,en;q=0.9,zh;q=0.8") + // Take the first language + parts := strings.Split(acceptLang, ",") + if len(parts) > 0 { + // Remove quality value if present + lang := strings.Split(parts[0], ";")[0] + return strings.ToLower(strings.TrimSpace(lang)) + } + } + + return "" +} + +// GetTheme extracts theme from request with priority: +// 1. Query parameter "theme" +// 2. Header "X-Yao-Theme" +func GetTheme(c *gin.Context) string { + // Priority 1: Query parameter + if theme := c.Query("theme"); theme != "" { + return strings.ToLower(theme) + } + + // Priority 2: Header + if theme := c.GetHeader("X-Yao-Theme"); theme != "" { + return strings.ToLower(theme) + } + + return "" +} + +// GetReferer extracts referer from request with priority: +// 1. Query parameter "referer" +// 2. Header "X-Yao-Referer" +// 3. Default to "api" +func GetReferer(c *gin.Context) string { + // Priority 1: Query parameter + if referer := c.Query("referer"); referer != "" { + return validateReferer(referer) + } + + // Priority 2: Header + if referer := c.GetHeader("X-Yao-Referer"); referer != "" { + return validateReferer(referer) + } + + // Priority 3: Default + return RefererAPI +} + +// GetAccept extracts accept type from request with priority: +// 1. Query parameter "accept" +// 2. Header "X-Yao-Accept" +// 3. Parse from client type (User-Agent) +func GetAccept(c *gin.Context) Accept { + // Priority 1: Query parameter + if accept := c.Query("accept"); accept != "" { + return validateAccept(accept) + } + + // Priority 2: Header + if accept := c.GetHeader("X-Yao-Accept"); accept != "" { + return validateAccept(accept) + } + + // Priority 3: Parse from User-Agent + userAgent := c.GetHeader("User-Agent") + clientType := getClientType(userAgent) + return parseAccept(clientType) +} + +// GetChatID get the chat ID from the request +// Priority: +// 1. Query parameter "chat_id" +// 2. Header "X-Yao-Chat" +// 3. Generate from messages using GetChatIDByMessages +func GetChatID(c *gin.Context, cache store.Store) (string, error) { + // Priority 1: Query parameter chat_id + if chatID := c.Query("chat_id"); chatID != "" { + return chatID, nil + } + + // Priority 2: Header X-Yao-Chat + if chatID := c.GetHeader("X-Yao-Chat"); chatID != "" { + return chatID, nil + } + + // Priority 3: Generate from messages + messages, err := GetMessages(c) + if err != nil { + return "", fmt.Errorf("failed to get messages for chat ID generation: %w", err) + } + + chatID, err := GetChatIDByMessages(cache, messages) + if err != nil { + return "", fmt.Errorf("failed to generate chat ID from messages: %w", err) + } + + return chatID, nil +} diff --git a/agent/context/openapi_chat_test.go b/agent/context/openapi_chat_test.go new file mode 100644 index 00000000..acfedb6e --- /dev/null +++ b/agent/context/openapi_chat_test.go @@ -0,0 +1,586 @@ +package context + +import ( + "bytes" + "encoding/json" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/gou/store" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/test" +) + +func TestGetMessages_FromBody(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + messages := []Message{ + { + Role: RoleUser, + Content: "Hello, world!", + }, + { + Role: RoleAssistant, + Content: "Hi there!", + }, + } + + requestBody := map[string]interface{}{ + "messages": messages, + "model": "gpt-4", + } + + bodyBytes, _ := json.Marshal(requestBody) + + req := httptest.NewRequest("POST", "/chat/completions", bytes.NewBuffer(bodyBytes)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + result, err := GetMessages(c) + if err != nil { + t.Fatalf("Failed to get messages: %v", err) + } + + if len(result) != 2 { + t.Errorf("Expected 2 messages, got %d", len(result)) + } + + if result[0].Role != RoleUser { + t.Errorf("Expected first message role to be %s, got %s", RoleUser, result[0].Role) + } +} + +func TestGetMessages_FromQuery(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + messages := []Message{ + { + Role: RoleUser, + Content: "Test message", + }, + } + + messagesJSON, _ := json.Marshal(messages) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + q := req.URL.Query() + q.Add("messages", string(messagesJSON)) + req.URL.RawQuery = q.Encode() + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + result, err := GetMessages(c) + if err != nil { + t.Fatalf("Failed to get messages: %v", err) + } + + if len(result) != 1 { + t.Errorf("Expected 1 message, got %d", len(result)) + } +} + +func TestGetMessages_EmptyMessages(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + requestBody := map[string]interface{}{ + "messages": []Message{}, + "model": "gpt-4", + } + + bodyBytes, _ := json.Marshal(requestBody) + + req := httptest.NewRequest("POST", "/chat/completions", bytes.NewBuffer(bodyBytes)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + _, err := GetMessages(c) + if err == nil { + t.Error("Expected error for empty messages") + } +} + +func TestGetChatID_FromQuery(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + cache, err := store.Get("__yao.agent.cache") + if err != nil { + t.Fatalf("Failed to get cache: %v", err) + } + + expectedChatID := "test-chat-123" + + req := httptest.NewRequest("GET", "/chat/completions?chat_id="+expectedChatID, nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + chatID, err := GetChatID(c, cache) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID != expectedChatID { + t.Errorf("Expected chat ID %s, got %s", expectedChatID, chatID) + } +} + +func TestGetChatID_FromHeader(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + cache, err := store.Get("__yao.agent.cache") + if err != nil { + t.Fatalf("Failed to get cache: %v", err) + } + + expectedChatID := "header-chat-456" + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("X-Yao-Chat", expectedChatID) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + chatID, err := GetChatID(c, cache) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID != expectedChatID { + t.Errorf("Expected chat ID %s, got %s", expectedChatID, chatID) + } +} + +func TestGetChatID_FromMessages(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + cache, err := store.Get("__yao.agent.cache") + if err != nil { + t.Fatalf("Failed to get cache: %v", err) + } + cache.Clear() + + // First request with one user message + messages1 := []Message{ + { + Role: RoleUser, + Content: "First message", + }, + } + + requestBody1 := map[string]interface{}{ + "messages": messages1, + } + + bodyBytes1, _ := json.Marshal(requestBody1) + + req := httptest.NewRequest("POST", "/chat/completions", bytes.NewBuffer(bodyBytes1)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + chatID1, err := GetChatID(c, cache) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID1 == "" { + t.Error("Expected non-empty chat ID") + } + + // Second request with two user messages (continuation) + messages2 := []Message{ + { + Role: RoleUser, + Content: "First message", + }, + { + Role: RoleUser, + Content: "Second message", + }, + } + + requestBody2 := map[string]interface{}{ + "messages": messages2, + } + + bodyBytes2, _ := json.Marshal(requestBody2) + + req2 := httptest.NewRequest("POST", "/chat/completions", bytes.NewBuffer(bodyBytes2)) + req2.Header.Set("Content-Type", "application/json") + w2 := httptest.NewRecorder() + c2, _ := gin.CreateTestContext(w2) + c2.Request = req2 + + chatID2, err := GetChatID(c2, cache) + if err != nil { + t.Fatalf("Failed to get chat ID second time: %v", err) + } + + // Should get same chat ID (continuation of conversation) + if chatID1 != chatID2 { + t.Errorf("Expected same chat ID for continuation, got %s and %s", chatID1, chatID2) + } +} + +func TestGetChatID_Priority(t *testing.T) { + test.Prepare(t, config.Conf) + defer test.Clean() + + gin.SetMode(gin.TestMode) + + cache, err := store.Get("__yao.agent.cache") + if err != nil { + t.Fatalf("Failed to get cache: %v", err) + } + + queryChatID := "query-chat-id" + headerChatID := "header-chat-id" + + messages := []Message{ + { + Role: RoleUser, + Content: "This should not be used", + }, + } + + requestBody := map[string]interface{}{ + "messages": messages, + } + + bodyBytes, _ := json.Marshal(requestBody) + + // Test priority: query > header > messages + req := httptest.NewRequest("POST", "/chat/completions?chat_id="+queryChatID, bytes.NewBuffer(bodyBytes)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Yao-Chat", headerChatID) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + chatID, err := GetChatID(c, cache) + if err != nil { + t.Fatalf("Failed to get chat ID: %v", err) + } + + if chatID != queryChatID { + t.Errorf("Expected query parameter to take priority, got %s instead of %s", chatID, queryChatID) + } +} + +func TestGetLocale_FromQuery(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?locale=zh-CN", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + locale := GetLocale(c) + if locale != "zh-cn" { + t.Errorf("Expected locale 'zh-cn', got '%s'", locale) + } +} + +func TestGetLocale_FromHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("Accept-Language", "en-US,en;q=0.9,zh;q=0.8") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + locale := GetLocale(c) + if locale != "en-us" { + t.Errorf("Expected locale 'en-us', got '%s'", locale) + } +} + +func TestGetLocale_Priority(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?locale=fr-FR", nil) + req.Header.Set("Accept-Language", "en-US") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + locale := GetLocale(c) + if locale != "fr-fr" { + t.Errorf("Expected query parameter to take priority, got '%s'", locale) + } +} + +func TestGetLocale_Empty(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + locale := GetLocale(c) + if locale != "" { + t.Errorf("Expected empty locale, got '%s'", locale) + } +} + +func TestGetTheme_FromQuery(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?theme=dark", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + theme := GetTheme(c) + if theme != "dark" { + t.Errorf("Expected theme 'dark', got '%s'", theme) + } +} + +func TestGetTheme_FromHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("X-Yao-Theme", "light") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + theme := GetTheme(c) + if theme != "light" { + t.Errorf("Expected theme 'light', got '%s'", theme) + } +} + +func TestGetTheme_Priority(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?theme=auto", nil) + req.Header.Set("X-Yao-Theme", "dark") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + theme := GetTheme(c) + if theme != "auto" { + t.Errorf("Expected query parameter to take priority, got '%s'", theme) + } +} + +func TestGetTheme_Empty(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + theme := GetTheme(c) + if theme != "" { + t.Errorf("Expected empty theme, got '%s'", theme) + } +} + +func TestGetReferer_FromQuery(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?referer=jssdk", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + referer := GetReferer(c) + if referer != "jssdk" { + t.Errorf("Expected referer 'jssdk', got '%s'", referer) + } +} + +func TestGetReferer_FromHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("X-Yao-Referer", "agent") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + referer := GetReferer(c) + if referer != "agent" { + t.Errorf("Expected referer 'agent', got '%s'", referer) + } +} + +func TestGetReferer_Default(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + referer := GetReferer(c) + if referer != RefererAPI { + t.Errorf("Expected default referer '%s', got '%s'", RefererAPI, referer) + } +} + +func TestGetReferer_InvalidValue(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?referer=invalid", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + referer := GetReferer(c) + if referer != RefererAPI { + t.Errorf("Expected default referer for invalid value, got '%s'", referer) + } +} + +func TestGetReferer_Priority(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?referer=process", nil) + req.Header.Set("X-Yao-Referer", "tool") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + referer := GetReferer(c) + if referer != "process" { + t.Errorf("Expected query parameter to take priority, got '%s'", referer) + } +} + +func TestGetAccept_FromQuery(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?accept=cui-web", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AcceptWebCUI { + t.Errorf("Expected accept '%s', got '%s'", AcceptWebCUI, accept) + } +} + +func TestGetAccept_FromHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("X-Yao-Accept", "cui-native") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AccepNativeCUI { + t.Errorf("Expected accept '%s', got '%s'", AccepNativeCUI, accept) + } +} + +func TestGetAccept_FromUserAgent_Web(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("User-Agent", "Mozilla/5.0") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AcceptWebCUI { + t.Errorf("Expected accept '%s' for web user agent, got '%s'", AcceptWebCUI, accept) + } +} + +func TestGetAccept_FromUserAgent_Android(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("User-Agent", "Android App") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AccepNativeCUI { + t.Errorf("Expected accept '%s' for Android, got '%s'", AccepNativeCUI, accept) + } +} + +func TestGetAccept_FromUserAgent_Desktop(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions", nil) + req.Header.Set("User-Agent", "Windows NT 10.0") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AcceptDesktopCUI { + t.Errorf("Expected accept '%s' for Windows, got '%s'", AcceptDesktopCUI, accept) + } +} + +func TestGetAccept_Priority(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?accept=standard", nil) + req.Header.Set("X-Yao-Accept", "cui-web") + req.Header.Set("User-Agent", "Android App") + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AcceptStandard { + t.Errorf("Expected query parameter to take priority, got '%s'", accept) + } +} + +func TestGetAccept_InvalidValue(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest("GET", "/chat/completions?accept=invalid", nil) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = req + + accept := GetAccept(c) + if accept != AcceptStandard { + t.Errorf("Expected default accept for invalid value, got '%s'", accept) + } +} diff --git a/agent/context/types.go b/agent/context/types.go index 0dca6750..df39a919 100644 --- a/agent/context/types.go +++ b/agent/context/types.go @@ -4,6 +4,7 @@ import ( "context" "github.com/yaoapp/gou/plan" + "github.com/yaoapp/gou/store" "github.com/yaoapp/yao/openapi/oauth/types" ) @@ -93,7 +94,8 @@ type Context struct { // Context context.Context - Space plan.Space `json:"-"` // Shared data space, it will be used to share data between the request and the call + Space plan.Space `json:"-"` // Shared data space, it will be used to share data between the request and the call + Cache store.Store `json:"-"` // Cache store, it will be used to store the message cache, default is "__yao.agent.cache" // Authorized information Authorized *types.AuthorizedInfo `json:"authorized,omitempty"` // Authorized information @@ -121,3 +123,96 @@ type Context struct { Silent bool `json:"silent,omitempty"` // Silent mode (Deprecated, use Referer instead) } + +// Message Structure ( OpenAI Chat Completion Input Message Structure, https://platform.openai.com/docs/api-reference/chat/create#chat/create-messages ) +// =============================== + +// MessageRole represents the role of a message author +type MessageRole string + +// Message role constants +const ( + RoleDeveloper MessageRole = "developer" // Developer-provided instructions (o1 models and newer) + RoleSystem MessageRole = "system" // System instructions + RoleUser MessageRole = "user" // User messages + RoleAssistant MessageRole = "assistant" // Assistant responses + RoleTool MessageRole = "tool" // Tool responses +) + +// Message represents a message in the conversation, compatible with OpenAI's chat completion API +// Supports message types: developer, system, user, assistant, and tool +type Message struct { + // Common fields for all message types + Role MessageRole `json:"role"` // Required: message author role + Content interface{} `json:"content,omitempty"` // string or array of ContentPart; Required for most types, optional for assistant with tool_calls + Name *string `json:"name,omitempty"` // Optional: participant name to differentiate between participants of the same role + + // Tool message specific fields + ToolCallID *string `json:"tool_call_id,omitempty"` // Required for tool messages: tool call that this message is responding to + + // Assistant message specific fields + ToolCalls []ToolCall `json:"tool_calls,omitempty"` // Optional for assistant: tool calls generated by the model + Refusal *string `json:"refusal,omitempty"` // Optional for assistant: refusal message (null when not refusing) +} + +// ContentPartType represents the type of content part +type ContentPartType string + +// Content part type constants +const ( + ContentText ContentPartType = "text" // Text content + ContentImageURL ContentPartType = "image_url" // Image URL content (Vision) + ContentInputAudio ContentPartType = "input_audio" // Input audio content (Audio) +) + +// ContentPart represents a part of the message content (for multimodal messages) +// Used when Content is an array instead of a simple string +type ContentPart struct { + Type ContentPartType `json:"type"` // Required: content part type + Text string `json:"text,omitempty"` // For type="text": the text content + ImageURL *ImageURL `json:"image_url,omitempty"` // For type="image_url": the image URL + InputAudio *InputAudio `json:"input_audio,omitempty"` // For type="input_audio": the input audio data +} + +// ImageDetailLevel represents the detail level for image processing +type ImageDetailLevel string + +// Image detail level constants +const ( + DetailAuto ImageDetailLevel = "auto" // Let the model decide + DetailLow ImageDetailLevel = "low" // Low detail (faster, cheaper) + DetailHigh ImageDetailLevel = "high" // High detail (slower, more expensive) +) + +// ImageURL represents an image URL in the message content +type ImageURL struct { + URL string `json:"url"` // Required: URL of the image or base64 encoded image data + Detail ImageDetailLevel `json:"detail,omitempty"` // Optional: how the model processes the image +} + +// InputAudio represents input audio data in the message content +type InputAudio struct { + Data string `json:"data"` // Required: Base64 encoded audio data + Format string `json:"format"` // Required: Audio format (e.g., "wav", "mp3") +} + +// ToolCallType represents the type of tool call +type ToolCallType string + +// Tool call type constants +const ( + ToolTypeFunction ToolCallType = "function" // Function call +) + +// ToolCall represents a tool call generated by the model (for assistant messages) +type ToolCall struct { + ID string `json:"id"` // Required: unique identifier for the tool call + Type ToolCallType `json:"type"` // Required: type of tool call, currently only "function" + Function Function `json:"function"` // Required: function call details +} + +// Function represents a function call with name and arguments +type Function struct { + Name string `json:"name"` // Required: name of the function to call + Arguments string `json:"arguments,omitempty"` // Optional: arguments to pass to the function, as a JSON string +} diff --git a/agent/context/utils.go b/agent/context/utils.go index 7b7c2621..5949c788 100644 --- a/agent/context/utils.go +++ b/agent/context/utils.go @@ -51,4 +51,3 @@ func parseAccept(clientType string) Accept { return AcceptStandard } } - diff --git a/agent/load.go b/agent/load.go index f5bc5b85..c3a4313a 100644 --- a/agent/load.go +++ b/agent/load.go @@ -6,23 +6,23 @@ import ( "github.com/yaoapp/gou/application" "github.com/yaoapp/gou/connector" + "github.com/yaoapp/kun/exception" + "github.com/yaoapp/yao/agent/api" "github.com/yaoapp/yao/agent/assistant" "github.com/yaoapp/yao/agent/i18n" mongoStore "github.com/yaoapp/yao/agent/store/mongo" redisStore "github.com/yaoapp/yao/agent/store/redis" store "github.com/yaoapp/yao/agent/store/types" xunStore "github.com/yaoapp/yao/agent/store/xun" + "github.com/yaoapp/yao/agent/types" "github.com/yaoapp/yao/config" ) -// Agent the agent AI assistant -var Agent *DSL - // Load load AIGC func Load(cfg config.Config) error { - setting := DSL{ - ID: "agent", + setting := types.DSL{ + Cache: "__yao.agent.cache", // default is "__yao.agent.cache" StoreSetting: store.Setting{ MaxSize: 20, TTL: 90 * 24 * 60 * 60, // 90 days in seconds @@ -45,7 +45,7 @@ func Load(cfg config.Config) error { // Default Assistant, Agent is the developer name, Mohe is the brand name of the assistant if setting.Use == nil { - setting.Use = &Use{Default: "mohe"} // Agent is the developer name, Mohe is the brand name of the assistant + setting.Use = &types.Use{Default: "mohe"} // Agent is the developer name, Mohe is the brand name of the assistant } // Title Assistant @@ -58,7 +58,8 @@ func Load(cfg config.Config) error { setting.Use.Prompt = setting.Use.Default } - Agent = &setting + // Initialize Agent API + api.Agent = &api.API{DSL: &setting} // Store Setting err = initStore() @@ -87,6 +88,14 @@ func Load(cfg config.Config) error { return nil } +// GetAgent returns the Agent instance +func GetAgent() *api.API { + if api.Agent == nil { + exception.New("Agent is not initialized", 500).Throw() + } + return api.Agent +} + // initGlobalI18n initialize the global i18n func initGlobalI18n() error { locales, err := i18n.GetLocales("agent") @@ -116,7 +125,7 @@ func initConnectorSettings() error { return err } - Agent.Connectors = connectors + api.Agent.DSL.Connectors = connectors return nil } @@ -124,46 +133,46 @@ func initConnectorSettings() error { func initStore() error { var err error - if Agent.StoreSetting.Connector == "default" || Agent.StoreSetting.Connector == "" { - Agent.Store, err = xunStore.NewXun(Agent.StoreSetting) + if api.Agent.DSL.StoreSetting.Connector == "default" || api.Agent.DSL.StoreSetting.Connector == "" { + api.Agent.DSL.Store, err = xunStore.NewXun(api.Agent.DSL.StoreSetting) return err } // other connector - conn, err := connector.Select(Agent.StoreSetting.Connector) + conn, err := connector.Select(api.Agent.DSL.StoreSetting.Connector) if err != nil { return fmt.Errorf("load connectors error: %s", err.Error()) } if conn.Is(connector.DATABASE) { - Agent.Store, err = xunStore.NewXun(Agent.StoreSetting) + api.Agent.DSL.Store, err = xunStore.NewXun(api.Agent.DSL.StoreSetting) return err } else if conn.Is(connector.REDIS) { - Agent.Store = redisStore.NewRedis() + api.Agent.DSL.Store = redisStore.NewRedis() return nil } else if conn.Is(connector.MONGO) { - Agent.Store = mongoStore.NewMongo() + api.Agent.DSL.Store = mongoStore.NewMongo() return nil } - return fmt.Errorf("%s store connector %s not support", Agent.ID, Agent.StoreSetting.Connector) + return fmt.Errorf("Agent store connector %s not support", api.Agent.DSL.StoreSetting.Connector) } // initAssistant initialize the assistant func initAssistant() error { // Set Storage - assistant.SetStorage(Agent.Store) + assistant.SetStorage(api.Agent.DSL.Store) // Assistant Vision - if Agent.Vision != nil { - assistant.SetVision(Agent.Vision) + if api.Agent.DSL.Vision != nil { + assistant.SetVision(api.Agent.DSL.Vision) } - if Agent.Connectors != nil { - assistant.SetConnectorSettings(Agent.Connectors) + if api.Agent.DSL.Connectors != nil { + assistant.SetConnectorSettings(api.Agent.DSL.Connectors) } // Load Built-in Assistants @@ -178,14 +187,14 @@ func initAssistant() error { return err } - Agent.Assistant = defaultAssistant + api.Agent.DSL.Assistant = defaultAssistant return nil } // defaultAssistant get the default assistant func defaultAssistant() (*assistant.Assistant, error) { - if Agent.Use == nil || Agent.Use.Default == "" { + if api.Agent.DSL.Use == nil || api.Agent.DSL.Use.Default == "" { return nil, fmt.Errorf("default assistant not found") } - return assistant.Get(Agent.Use.Default) + return assistant.Get(api.Agent.DSL.Use.Default) } diff --git a/agent/process.go b/agent/process.go index d344b18b..7ce6d26c 100644 --- a/agent/process.go +++ b/agent/process.go @@ -12,14 +12,6 @@ import ( store "github.com/yaoapp/yao/agent/store/types" ) -// GetAgent returns the Agent instance -func GetAgent() *DSL { - if Agent == nil { - exception.New("Agent is not initialized", 500).Throw() - } - return Agent -} - func init() { process.RegisterGroup("agent", map[string]process.Handler{ "write": ProcessWrite, @@ -201,11 +193,11 @@ func parseAssistantFilter(params map[string]interface{}) store.AssistantFilter { case []interface{}: filter.Select = []string{} for _, field := range v { - switch field.(type) { + switch v := field.(type) { case string: - filter.Select = append(filter.Select, field.(string)) + filter.Select = append(filter.Select, v) case interface{}: - filter.Select = append(filter.Select, fmt.Sprintf("%v", field)) + filter.Select = append(filter.Select, fmt.Sprintf("%v", v)) } } diff --git a/agent/types/dsl.go b/agent/types/dsl.go new file mode 100644 index 00000000..85b2e3ac --- /dev/null +++ b/agent/types/dsl.go @@ -0,0 +1,11 @@ +package types + +import "github.com/yaoapp/gou/store" + +// GetCacheStore get the cache store +func (dsl *DSL) GetCacheStore() (store.Store, error) { + if dsl.Cache == "" { + return store.Get("__yao.agent.cache") + } + return store.Get(dsl.Cache) +} diff --git a/agent/types.go b/agent/types/types.go similarity index 95% rename from agent/types.go rename to agent/types/types.go index d59936ef..446d4a59 100644 --- a/agent/types.go +++ b/agent/types/types.go @@ -1,4 +1,4 @@ -package agent +package types import ( "github.com/gin-gonic/gin" @@ -14,6 +14,8 @@ type DSL struct { // =============================== Use *Use `json:"use,omitempty" yaml:"use,omitempty"` // Which assistant to use default, title, prompt StoreSetting store.Setting `json:"store" yaml:"store"` // The store setting of the assistant + Cache string `json:"cache" yaml:"cache"` // The cache store of the assistant, if not set, default is "__yao.agent.cache" + // AuthSetting *Auth `json:"auth,omitempty" yaml:"auth,omitempty"` // Authenticate Settings // UploadSetting *Upload `json:"upload,omitempty" yaml:"upload,omitempty"` // Upload Settings // KnowledgeSetting *Knowledge `json:"knowledge,omitempty" yaml:"knowledge,omitempty"` // Knowledge base Settings @@ -29,7 +31,7 @@ type DSL struct { // Internal // =============================== - ID string `json:"-" yaml:"-"` // The id of the instance + // ID string `json:"-" yaml:"-"` // The id of the instance Assistant assistant.API `json:"-" yaml:"-"` // The default assistant Store store.Store `json:"-" yaml:"-"` // The store of the assistant Vision *vision.Vision `json:"-" yaml:"-"` diff --git a/openapi/agent/agent.go b/openapi/agent/agent.go index db79fd9f..307f0745 100644 --- a/openapi/agent/agent.go +++ b/openapi/agent/agent.go @@ -2,7 +2,6 @@ package agent import ( "github.com/gin-gonic/gin" - "github.com/yaoapp/yao/agent" "github.com/yaoapp/yao/openapi/oauth/types" ) @@ -11,19 +10,19 @@ import ( func Attach(group *gin.RouterGroup, oauth types.OAuth) { // Get the Agent instance - n := agent.GetAgent() + // n := agent.GetAgent() // Apply OAuth guard to all routes group.Use(oauth.Guard) // Assistant CRUD - Standard REST endpoints - group.GET("/assistants", ListAssistants) // GET /assistants - List assistants - group.POST("/assistants", CreateAssistant) // POST /assistants - Create assistant - group.GET("/assistants/tags", ListAssistantTags) // GET /assistants/tags - Get all assistant tags with permission filtering - group.GET("/assistants/:id", GetAssistant) // GET /assistants/:id - Get assistant details with permission verification - group.PUT("/assistants/:id", UpdateAssistant) // PUT /assistants/:id - Update assistant - group.DELETE("/assistants/:id", n.HandleAssistantDelete) // DELETE /assistants/:id - Delete assistant + group.GET("/assistants", ListAssistants) // GET /assistants - List assistants + group.POST("/assistants", CreateAssistant) // POST /assistants - Create assistant + group.GET("/assistants/tags", ListAssistantTags) // GET /assistants/tags - Get all assistant tags with permission filtering + group.GET("/assistants/:id", GetAssistant) // GET /assistants/:id - Get assistant details with permission verification + group.PUT("/assistants/:id", UpdateAssistant) // PUT /assistants/:id - Update assistant + // group.DELETE("/assistants/:id", agent.HandleAssistantDelete) // DELETE /assistants/:id - Delete assistant // Assistant Actions - group.POST("/assistants/:id/call", n.HandleAssistantCall) // POST /assistants/:id/call - Execute assistant API + // group.POST("/assistants/:id/call", agent.HandleAssistantCall) // POST /assistants/:id/call - Execute assistant API } diff --git a/openapi/chat/completions.go b/openapi/chat/completions.go index 25947d27..66d39e80 100644 --- a/openapi/chat/completions.go +++ b/openapi/chat/completions.go @@ -9,16 +9,45 @@ import ( "github.com/gin-gonic/gin" "github.com/google/uuid" + "github.com/yaoapp/yao/agent" + "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/openapi/response" ) // GinCreateCompletions handles POST /chat/:assistant_id/completions - Create a chat completion func GinCreateCompletions(c *gin.Context) { - // Print request information for debugging - fmt.Println("========== Chat Completions Request ==========") - fmt.Printf("Method: %s\n", c.Request.Method) - fmt.Printf("URL: %s\n", c.Request.URL.String()) - fmt.Printf("RemoteAddr: %s\n", c.Request.RemoteAddr) + + agent := agent.GetAgent() + cache, err := agent.GetCacheStore() + if err != nil { + response.RespondWithError(c, response.StatusInternalServerError, &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get cache store: " + err.Error(), + }) + return + } + + ctx := context.NewOpenAPI(c, cache) + + fmt.Println("-----------------------------------------------") + fmt.Println("Chat ID: ", ctx.ChatID) + fmt.Println("Assistant ID: ", ctx.AssistantID) + fmt.Println("----") + // utils.Dump(ctx) + // Print request body + fmt.Println("\n--- Request Body ---") + body, err := io.ReadAll(c.Request.Body) + if err != nil { + fmt.Printf("Error reading body: %v\n", err) + } else { + fmt.Printf("%s\n", string(body)) + // Restore the body for further processing + c.Request.Body = io.NopCloser(bytes.NewBuffer(body)) + } + fmt.Println("-----------------------------------------------") + + c.JSON(response.StatusOK, gin.H{"message": "Create Completions", "chat_id": ctx.ChatID}) + return // Print headers fmt.Println("\n--- Headers ---") @@ -44,7 +73,7 @@ func GinCreateCompletions(c *gin.Context) { // Print request body fmt.Println("\n--- Request Body ---") - body, err := io.ReadAll(c.Request.Body) + body, err = io.ReadAll(c.Request.Body) if err != nil { fmt.Printf("Error reading body: %v\n", err) } else { diff --git a/openapi/oauth/guard.go b/openapi/oauth/guard.go index 7230e798..6779d2a4 100644 --- a/openapi/oauth/guard.go +++ b/openapi/oauth/guard.go @@ -146,10 +146,10 @@ func (s *Service) getAccessTokenFromAPIKey(apiKey string) string { log.Warn("Failed to make access token: %s", err.Error()) } - fmt.Println("========== Access Token From API Key ==========") - fmt.Println("accessToken: ", accessToken) - fmt.Println("extraClaims: ", extraClaims) - fmt.Println("===============================================") + // fmt.Println("========== Access Token From API Key ==========") + // fmt.Println("accessToken: ", accessToken) + // fmt.Println("extraClaims: ", extraClaims) + // fmt.Println("===============================================") return accessToken } diff --git a/service/service.go b/service/service.go index e1d15139..0888144a 100644 --- a/service/service.go +++ b/service/service.go @@ -6,7 +6,7 @@ import ( "github.com/gin-gonic/gin" "github.com/yaoapp/gou/api" "github.com/yaoapp/gou/server/http" - "github.com/yaoapp/yao/agent" + agent "github.com/yaoapp/yao/agent/api" "github.com/yaoapp/yao/config" "github.com/yaoapp/yao/openapi" "github.com/yaoapp/yao/share" diff --git a/test/utils.go b/test/utils.go index 79c3fc95..090f3783 100644 --- a/test/utils.go +++ b/test/utils.go @@ -223,6 +223,7 @@ var testSystemStores = map[string]string{ "__yao.oauth.client": "yao/stores/oauth/client.badger.yao", "__yao.oauth.cache": "yao/stores/oauth/cache.lru.yao", "__yao.agent.memory": "yao/stores/agent/memory.badger.yao", + "__yao.agent.cache": "yao/stores/agent/cache.lru.yao", "__yao.kb.store": "yao/stores/kb/store.badger.yao", "__yao.kb.cache": "yao/stores/kb/cache.lru.yao", } diff --git a/widgets/app/app.go b/widgets/app/app.go index bd1d8035..811feb0b 100644 --- a/widgets/app/app.go +++ b/widgets/app/app.go @@ -18,7 +18,7 @@ import ( "github.com/yaoapp/gou/session" "github.com/yaoapp/kun/exception" "github.com/yaoapp/kun/log" - "github.com/yaoapp/yao/agent" + agent "github.com/yaoapp/yao/agent/api" "github.com/yaoapp/yao/agent/assistant" "github.com/yaoapp/yao/config" "github.com/yaoapp/yao/data" @@ -47,7 +47,7 @@ import ( // Setting the application setting var Setting *DSL -var regExcp = regexp.MustCompile("^Exception\\|([0-9]+):(.+)$") +var regExcp = regexp.MustCompile(`^Exception\|([0-9]+):(.+)$`) // LoadAndExport load app func LoadAndExport(cfg config.Config) error { @@ -258,7 +258,7 @@ func processService(process *process.Process) interface{} { process.ValidateArgNums(2) service := fmt.Sprintf("__yao_service.%s", process.ArgsString(0)) payload := process.ArgsMap(1) - if payload == nil || len(payload) == 0 { + if len(payload) == 0 { exception.New("content is required", 400).Throw() } @@ -508,7 +508,7 @@ func processXgen(process *process.Process) interface{} { "layout": layout, } - if admin.ThirdPartyLogin != nil && len(admin.ThirdPartyLogin) > 0 { + if len(admin.ThirdPartyLogin) > 0 { xgenLogin["admin"]["thirdPartyLogin"] = admin.ThirdPartyLogin } @@ -544,7 +544,7 @@ func processXgen(process *process.Process) interface{} { "layout": layout, } - if user.ThirdPartyLogin != nil && len(user.ThirdPartyLogin) > 0 { + if len(user.ThirdPartyLogin) > 0 { xgenLogin["user"]["thirdPartyLogin"] = user.ThirdPartyLogin } } @@ -715,8 +715,9 @@ func (dsl *DSL) replaceAdminRoot() error { // icons func (dsl *DSL) icons(cfg config.Config) { - dsl.Favicon = fmt.Sprintf("/api/__yao/app/icons/app.ico") - dsl.Logo = fmt.Sprintf("/api/__yao/app/icons/app.png") + dsl.Favicon = "/api/__yao/app/icons/app.ico" + dsl.Logo = "/api/__yao/app/icons/app.png" + log.Trace("CFG %v", cfg.Root) // favicon := filepath.Join(cfg.Root, "icons", "app.ico") // if _, err := os.Stat(favicon); err == nil { @@ -745,23 +746,19 @@ func Permissions(process *process.Process, widget string, id string) map[string] for _, value := range values { permissions[fmt.Sprintf("%v", value)] = true } - break case []string: for _, value := range values { permissions[value] = true } - break case map[string]interface{}: for key := range values { permissions[key] = true } - break case map[string]bool: permissions = values - break } return permissions