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.
This commit is contained in:
parent
1c502bfea8
commit
8dca4719c0
25 changed files with 2197 additions and 200 deletions
|
|
@ -1,13 +1,22 @@
|
||||||
package agent
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/yaoapp/yao/agent/assistant"
|
"github.com/yaoapp/yao/agent/assistant"
|
||||||
chatctx "github.com/yaoapp/yao/agent/context"
|
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
|
// 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 err error
|
||||||
var ast assistant.API = Agent.Assistant
|
var ast assistant.API = Agent.Assistant
|
||||||
if ctx.AssistantID != "" {
|
if ctx.AssistantID != "" {
|
||||||
|
|
@ -21,7 +30,7 @@ func (agent *DSL) Answer(ctx chatctx.Context, question string, c *gin.Context) e
|
||||||
}
|
}
|
||||||
|
|
||||||
// Select select an assistant
|
// Select select an assistant
|
||||||
func (agent *DSL) Select(id string) (assistant.API, error) {
|
func (agent *API) Select(id string) (assistant.API, error) {
|
||||||
if id == "" {
|
if id == "" {
|
||||||
return Agent.Assistant, nil
|
return Agent.Assistant, nil
|
||||||
}
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package agent
|
package api
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
@ -14,12 +14,13 @@ import (
|
||||||
chatctx "github.com/yaoapp/yao/agent/context"
|
chatctx "github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/message"
|
"github.com/yaoapp/yao/agent/message"
|
||||||
store "github.com/yaoapp/yao/agent/store/types"
|
store "github.com/yaoapp/yao/agent/store/types"
|
||||||
|
"github.com/yaoapp/yao/agent/types"
|
||||||
"github.com/yaoapp/yao/helper"
|
"github.com/yaoapp/yao/helper"
|
||||||
"github.com/yaoapp/yao/openapi/oauth"
|
"github.com/yaoapp/yao/openapi/oauth"
|
||||||
)
|
)
|
||||||
|
|
||||||
// API registers the Agent API endpoints
|
// 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
|
// Get the guards
|
||||||
middlewares, err := agent.getGuardHandlers()
|
middlewares, err := agent.getGuardHandlers()
|
||||||
|
|
@ -145,13 +146,13 @@ func (agent *DSL) API(router *gin.Engine, path string) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleStatus handles the status request
|
// handleStatus handles the status request
|
||||||
func (agent *DSL) handleStatus(c *gin.Context) {
|
func (agent *API) handleStatus(c *gin.Context) {
|
||||||
c.Status(200)
|
c.Status(200)
|
||||||
c.Done()
|
c.Done()
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleChat handles the chat request
|
// handleChat handles the chat request
|
||||||
func (agent *DSL) handleChat(c *gin.Context) {
|
func (agent *API) handleChat(c *gin.Context) {
|
||||||
// Set headers for SSE
|
// Set headers for SSE
|
||||||
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
||||||
c.Header("Cache-Control", "no-cache")
|
c.Header("Cache-Control", "no-cache")
|
||||||
|
|
@ -215,7 +216,7 @@ func (agent *DSL) handleChat(c *gin.Context) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleChatList handles the chat list request
|
// handleChatList handles the chat list request
|
||||||
func (agent *DSL) handleChatList(c *gin.Context) {
|
func (agent *API) handleChatList(c *gin.Context) {
|
||||||
sid := c.GetString("__sid")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// handleChatHistory handles the chat history request
|
||||||
func (agent *DSL) handleChatHistory(c *gin.Context) {
|
func (agent *API) handleChatHistory(c *gin.Context) {
|
||||||
sid := c.GetString("__sid")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// 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")
|
origin := c.Request.Header.Get("Origin")
|
||||||
if origin == "" {
|
if origin == "" {
|
||||||
origin = c.Request.Referer()
|
origin = c.Request.Referer()
|
||||||
|
|
@ -294,12 +295,12 @@ func (agent *DSL) getOrigin(c *gin.Context) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
// getGuardHandlers returns authentication middleware handlers
|
// getGuardHandlers returns authentication middleware handlers
|
||||||
func (agent *DSL) getGuardHandlers() ([]gin.HandlerFunc, error) {
|
func (agent *API) getGuardHandlers() ([]gin.HandlerFunc, error) {
|
||||||
return []gin.HandlerFunc{}, nil
|
return []gin.HandlerFunc{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// defaultGuard is the default authentication handler
|
// 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
|
// Check if the request is for OpenAPI OAuth
|
||||||
if oauth.OAuth != nil {
|
if oauth.OAuth != nil {
|
||||||
|
|
@ -320,7 +321,7 @@ func (agent *DSL) defaultGuard(c *gin.Context) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Openapi Oauth
|
// Openapi Oauth
|
||||||
func (agent *DSL) guardOpenapiOauth(c *gin.Context) {
|
func (agent *API) guardOpenapiOauth(c *gin.Context) {
|
||||||
s := oauth.OAuth
|
s := oauth.OAuth
|
||||||
token := agent.getAccessToken(c)
|
token := agent.getAccessToken(c)
|
||||||
if token == "" {
|
if token == "" {
|
||||||
|
|
@ -348,7 +349,7 @@ func (agent *DSL) guardOpenapiOauth(c *gin.Context) {
|
||||||
c.Set("__sid", sid)
|
c.Set("__sid", sid)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (agent *DSL) getAccessToken(c *gin.Context) string {
|
func (agent *API) getAccessToken(c *gin.Context) string {
|
||||||
token := c.GetHeader("Authorization")
|
token := c.GetHeader("Authorization")
|
||||||
if token == "" || token == "Bearer undefined" {
|
if token == "" || token == "Bearer undefined" {
|
||||||
cookie, err := c.Cookie("__Host-access_token")
|
cookie, err := c.Cookie("__Host-access_token")
|
||||||
|
|
@ -360,7 +361,7 @@ func (agent *DSL) getAccessToken(c *gin.Context) string {
|
||||||
return strings.TrimPrefix(token, "Bearer ")
|
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")
|
sid, err := c.Cookie("__Host-session_id")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ""
|
return ""
|
||||||
|
|
@ -369,7 +370,7 @@ func (agent *DSL) getSessionID(c *gin.Context) string {
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleChatLatest handles getting the latest chat
|
// handleChatLatest handles getting the latest chat
|
||||||
func (agent *DSL) handleChatLatest(c *gin.Context) {
|
func (agent *API) handleChatLatest(c *gin.Context) {
|
||||||
sid := c.GetString("__sid")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// 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")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// 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")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// Convert assistants to mentions
|
||||||
mentions := []Mention{}
|
mentions := []types.Mention{}
|
||||||
for _, assistant := range response.Data {
|
for _, assistant := range response.Data {
|
||||||
mention := Mention{
|
mention := types.Mention{
|
||||||
ID: assistant.ID,
|
ID: assistant.ID,
|
||||||
Name: assistant.Name,
|
Name: assistant.Name,
|
||||||
Type: assistant.Type,
|
Type: assistant.Type,
|
||||||
|
|
@ -552,7 +553,7 @@ func (agent *DSL) handleMentions(c *gin.Context) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleChatUpdate handles updating a chat's details
|
// 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")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// handleChatDelete handles deleting a single chat
|
||||||
func (agent *DSL) handleChatDelete(c *gin.Context) {
|
func (agent *API) handleChatDelete(c *gin.Context) {
|
||||||
sid := c.GetString("__sid")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// 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")
|
sid := c.GetString("__sid")
|
||||||
if sid == "" {
|
if sid == "" {
|
||||||
c.JSON(400, gin.H{"message": "sid is required", "code": 400})
|
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
|
// handleGenerateTitle handles generating a chat title
|
||||||
func (agent *DSL) handleGenerateTitle(c *gin.Context) {
|
func (agent *API) handleGenerateTitle(c *gin.Context) {
|
||||||
// Set headers for SSE
|
// Set headers for SSE
|
||||||
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
||||||
c.Header("Cache-Control", "no-cache")
|
c.Header("Cache-Control", "no-cache")
|
||||||
|
|
@ -683,7 +684,7 @@ func (agent *DSL) handleGenerateTitle(c *gin.Context) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleGeneratePrompts handles generating prompts
|
// handleGeneratePrompts handles generating prompts
|
||||||
func (agent *DSL) handleGeneratePrompts(c *gin.Context) {
|
func (agent *API) handleGeneratePrompts(c *gin.Context) {
|
||||||
// Set headers for SSE
|
// Set headers for SSE
|
||||||
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
c.Header("Content-Type", "text/event-stream;charset=utf-8")
|
||||||
c.Header("Cache-Control", "no-cache")
|
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)
|
// 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
|
// Parse filter parameters
|
||||||
filter := store.AssistantFilter{
|
filter := store.AssistantFilter{
|
||||||
Type: "assistant",
|
Type: "assistant",
|
||||||
|
|
@ -825,7 +826,7 @@ func parseBoolValue(value string) *bool {
|
||||||
}
|
}
|
||||||
|
|
||||||
// HandleAssistantCall handles the assistant API call (exported for use in openapi/agent)
|
// 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")
|
assistantID := c.Param("id")
|
||||||
if assistantID == "" {
|
if assistantID == "" {
|
||||||
c.JSON(400, gin.H{"message": "assistant id is required", "code": 400})
|
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)
|
// 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")
|
assistantID := c.Param("id")
|
||||||
if assistantID == "" {
|
if assistantID == "" {
|
||||||
c.JSON(400, gin.H{"message": "assistant id is required", "code": 400})
|
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)
|
// 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{}
|
var assistantData map[string]interface{}
|
||||||
if err := c.BindJSON(&assistantData); err != nil {
|
if err := c.BindJSON(&assistantData); err != nil {
|
||||||
c.JSON(400, gin.H{"message": "invalid request body", "code": 400})
|
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)
|
// 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")
|
assistantID := c.Param("id")
|
||||||
if assistantID == "" {
|
if assistantID == "" {
|
||||||
c.JSON(400, gin.H{"message": "assistant id is required", "code": 400})
|
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
|
// handleConnectors handles listing connectors
|
||||||
func (agent *DSL) handleConnectors(c *gin.Context) {
|
func (agent *API) handleConnectors(c *gin.Context) {
|
||||||
options := []map[string]interface{}{}
|
options := []map[string]interface{}{}
|
||||||
|
|
||||||
// Filter and format connectors
|
// 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)
|
// 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
|
locale := "en-us" // Default locale
|
||||||
if loc := c.Query("locale"); loc != "" {
|
if loc := c.Query("locale"); loc != "" {
|
||||||
locale = strings.ToLower(strings.TrimSpace(loc))
|
locale = strings.ToLower(strings.TrimSpace(loc))
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package agent
|
package api
|
||||||
|
|
||||||
// import (
|
// import (
|
||||||
// "context"
|
// "context"
|
||||||
|
|
@ -302,7 +302,7 @@ func (ast *Assistant) Call(c *gin.Context, payload APIPayload) (interface{}, err
|
||||||
return nil, fmt.Errorf(HookErrorMethodNotFound)
|
return nil, fmt.Errorf(HookErrorMethodNotFound)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.Args == nil || len(payload.Args) == 0 {
|
if len(payload.Args) == 0 {
|
||||||
return scriptCtx.CallWith(ctx, method)
|
return scriptCtx.CallWith(ctx, method)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
182
agent/context/chat.go
Normal file
182
agent/context/chat.go
Normal file
|
|
@ -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
|
||||||
|
}
|
||||||
375
agent/context/chat_test.go
Normal file
375
agent/context/chat_test.go
Normal file
|
|
@ -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.
|
||||||
|
}
|
||||||
|
|
@ -2,7 +2,6 @@ package context
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
jsoniter "github.com/json-iterator/go"
|
jsoniter "github.com/json-iterator/go"
|
||||||
|
|
@ -126,8 +125,3 @@ func (ctx *Context) Map() map[string]interface{} {
|
||||||
|
|
||||||
return data
|
return data
|
||||||
}
|
}
|
||||||
|
|
||||||
// GenChatID generate a new chat ID
|
|
||||||
func GenChatID() string {
|
|
||||||
return fmt.Sprintf("chat_%d", time.Now().UnixNano())
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,7 @@ import (
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewGin(t *testing.T) {
|
func TestNewOpenAPI(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
|
|
@ -28,21 +28,22 @@ func TestNewGin(t *testing.T) {
|
||||||
{
|
{
|
||||||
name: "Parse all query parameters",
|
name: "Parse all query parameters",
|
||||||
queryParams: map[string]string{
|
queryParams: map[string]string{
|
||||||
"chat_id": "chat123",
|
"assistant_id": "ast456",
|
||||||
"locale": "zh-CN",
|
"chat_id": "chat123",
|
||||||
"theme": "dark",
|
"locale": "zh-CN",
|
||||||
"referer": RefererProcess,
|
"theme": "Dark",
|
||||||
"accept": string(AcceptStandard),
|
"referer": RefererProcess,
|
||||||
|
"accept": string(AcceptStandard),
|
||||||
},
|
},
|
||||||
routeParams: map[string]string{
|
routeParams: map[string]string{
|
||||||
"assistant_id": "ast456",
|
"assistant_id": "route-ast",
|
||||||
},
|
},
|
||||||
headers: map[string]string{
|
headers: map[string]string{
|
||||||
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
||||||
},
|
},
|
||||||
expectedChatID: "chat123",
|
expectedChatID: "chat123",
|
||||||
expectedAssistant: "ast456",
|
expectedAssistant: "ast456",
|
||||||
expectedLocale: "zh-CN",
|
expectedLocale: "zh-cn",
|
||||||
expectedTheme: "dark",
|
expectedTheme: "dark",
|
||||||
expectedClientType: "macos",
|
expectedClientType: "macos",
|
||||||
expectedReferer: RefererProcess,
|
expectedReferer: RefererProcess,
|
||||||
|
|
@ -168,7 +169,7 @@ func TestNewGin(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Call NewGin
|
// Call NewGin
|
||||||
ctx := NewGin(c)
|
ctx := NewOpenAPI(c, nil)
|
||||||
|
|
||||||
// Assertions
|
// Assertions
|
||||||
assert.Equal(t, tt.expectedChatID, ctx.ChatID, "ChatID mismatch")
|
assert.Equal(t, tt.expectedChatID, ctx.ChatID, "ChatID mismatch")
|
||||||
|
|
@ -204,7 +205,7 @@ func TestParseClientType(t *testing.T) {
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
result := parseClientType(tt.userAgent)
|
result := getClientType(tt.userAgent)
|
||||||
assert.Equal(t, tt.expected, result)
|
assert.Equal(t, tt.expected, result)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
96
agent/context/message.go
Normal file
96
agent/context/message.go
Normal file
|
|
@ -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 != ""
|
||||||
|
}
|
||||||
387
agent/context/message_test.go
Normal file
387
agent/context/message_test.go
Normal file
|
|
@ -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))
|
||||||
|
}
|
||||||
|
}
|
||||||
313
agent/context/openapi.go
Normal file
313
agent/context/openapi.go
Normal file
|
|
@ -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
|
||||||
|
}
|
||||||
586
agent/context/openapi_chat_test.go
Normal file
586
agent/context/openapi_chat_test.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
"github.com/yaoapp/gou/plan"
|
"github.com/yaoapp/gou/plan"
|
||||||
|
"github.com/yaoapp/gou/store"
|
||||||
"github.com/yaoapp/yao/openapi/oauth/types"
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -93,7 +94,8 @@ type Context struct {
|
||||||
|
|
||||||
// Context
|
// Context
|
||||||
context.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 information
|
||||||
Authorized *types.AuthorizedInfo `json:"authorized,omitempty"` // 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)
|
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
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -51,4 +51,3 @@ func parseAccept(clientType string) Accept {
|
||||||
return AcceptStandard
|
return AcceptStandard
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -6,23 +6,23 @@ import (
|
||||||
|
|
||||||
"github.com/yaoapp/gou/application"
|
"github.com/yaoapp/gou/application"
|
||||||
"github.com/yaoapp/gou/connector"
|
"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/assistant"
|
||||||
"github.com/yaoapp/yao/agent/i18n"
|
"github.com/yaoapp/yao/agent/i18n"
|
||||||
mongoStore "github.com/yaoapp/yao/agent/store/mongo"
|
mongoStore "github.com/yaoapp/yao/agent/store/mongo"
|
||||||
redisStore "github.com/yaoapp/yao/agent/store/redis"
|
redisStore "github.com/yaoapp/yao/agent/store/redis"
|
||||||
store "github.com/yaoapp/yao/agent/store/types"
|
store "github.com/yaoapp/yao/agent/store/types"
|
||||||
xunStore "github.com/yaoapp/yao/agent/store/xun"
|
xunStore "github.com/yaoapp/yao/agent/store/xun"
|
||||||
|
"github.com/yaoapp/yao/agent/types"
|
||||||
"github.com/yaoapp/yao/config"
|
"github.com/yaoapp/yao/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Agent the agent AI assistant
|
|
||||||
var Agent *DSL
|
|
||||||
|
|
||||||
// Load load AIGC
|
// Load load AIGC
|
||||||
func Load(cfg config.Config) error {
|
func Load(cfg config.Config) error {
|
||||||
|
|
||||||
setting := DSL{
|
setting := types.DSL{
|
||||||
ID: "agent",
|
Cache: "__yao.agent.cache", // default is "__yao.agent.cache"
|
||||||
StoreSetting: store.Setting{
|
StoreSetting: store.Setting{
|
||||||
MaxSize: 20,
|
MaxSize: 20,
|
||||||
TTL: 90 * 24 * 60 * 60, // 90 days in seconds
|
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
|
// Default Assistant, Agent is the developer name, Mohe is the brand name of the assistant
|
||||||
if setting.Use == nil {
|
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
|
// Title Assistant
|
||||||
|
|
@ -58,7 +58,8 @@ func Load(cfg config.Config) error {
|
||||||
setting.Use.Prompt = setting.Use.Default
|
setting.Use.Prompt = setting.Use.Default
|
||||||
}
|
}
|
||||||
|
|
||||||
Agent = &setting
|
// Initialize Agent API
|
||||||
|
api.Agent = &api.API{DSL: &setting}
|
||||||
|
|
||||||
// Store Setting
|
// Store Setting
|
||||||
err = initStore()
|
err = initStore()
|
||||||
|
|
@ -87,6 +88,14 @@ func Load(cfg config.Config) error {
|
||||||
return nil
|
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
|
// initGlobalI18n initialize the global i18n
|
||||||
func initGlobalI18n() error {
|
func initGlobalI18n() error {
|
||||||
locales, err := i18n.GetLocales("agent")
|
locales, err := i18n.GetLocales("agent")
|
||||||
|
|
@ -116,7 +125,7 @@ func initConnectorSettings() error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
Agent.Connectors = connectors
|
api.Agent.DSL.Connectors = connectors
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -124,46 +133,46 @@ func initConnectorSettings() error {
|
||||||
func initStore() error {
|
func initStore() error {
|
||||||
|
|
||||||
var err error
|
var err error
|
||||||
if Agent.StoreSetting.Connector == "default" || Agent.StoreSetting.Connector == "" {
|
if api.Agent.DSL.StoreSetting.Connector == "default" || api.Agent.DSL.StoreSetting.Connector == "" {
|
||||||
Agent.Store, err = xunStore.NewXun(Agent.StoreSetting)
|
api.Agent.DSL.Store, err = xunStore.NewXun(api.Agent.DSL.StoreSetting)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// other connector
|
// other connector
|
||||||
conn, err := connector.Select(Agent.StoreSetting.Connector)
|
conn, err := connector.Select(api.Agent.DSL.StoreSetting.Connector)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("load connectors error: %s", err.Error())
|
return fmt.Errorf("load connectors error: %s", err.Error())
|
||||||
}
|
}
|
||||||
|
|
||||||
if conn.Is(connector.DATABASE) {
|
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
|
return err
|
||||||
|
|
||||||
} else if conn.Is(connector.REDIS) {
|
} else if conn.Is(connector.REDIS) {
|
||||||
Agent.Store = redisStore.NewRedis()
|
api.Agent.DSL.Store = redisStore.NewRedis()
|
||||||
return nil
|
return nil
|
||||||
|
|
||||||
} else if conn.Is(connector.MONGO) {
|
} else if conn.Is(connector.MONGO) {
|
||||||
Agent.Store = mongoStore.NewMongo()
|
api.Agent.DSL.Store = mongoStore.NewMongo()
|
||||||
return nil
|
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
|
// initAssistant initialize the assistant
|
||||||
func initAssistant() error {
|
func initAssistant() error {
|
||||||
|
|
||||||
// Set Storage
|
// Set Storage
|
||||||
assistant.SetStorage(Agent.Store)
|
assistant.SetStorage(api.Agent.DSL.Store)
|
||||||
|
|
||||||
// Assistant Vision
|
// Assistant Vision
|
||||||
if Agent.Vision != nil {
|
if api.Agent.DSL.Vision != nil {
|
||||||
assistant.SetVision(Agent.Vision)
|
assistant.SetVision(api.Agent.DSL.Vision)
|
||||||
}
|
}
|
||||||
|
|
||||||
if Agent.Connectors != nil {
|
if api.Agent.DSL.Connectors != nil {
|
||||||
assistant.SetConnectorSettings(Agent.Connectors)
|
assistant.SetConnectorSettings(api.Agent.DSL.Connectors)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load Built-in Assistants
|
// Load Built-in Assistants
|
||||||
|
|
@ -178,14 +187,14 @@ func initAssistant() error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
Agent.Assistant = defaultAssistant
|
api.Agent.DSL.Assistant = defaultAssistant
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// defaultAssistant get the default assistant
|
// defaultAssistant get the default assistant
|
||||||
func defaultAssistant() (*assistant.Assistant, error) {
|
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 nil, fmt.Errorf("default assistant not found")
|
||||||
}
|
}
|
||||||
return assistant.Get(Agent.Use.Default)
|
return assistant.Get(api.Agent.DSL.Use.Default)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,14 +12,6 @@ import (
|
||||||
store "github.com/yaoapp/yao/agent/store/types"
|
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() {
|
func init() {
|
||||||
process.RegisterGroup("agent", map[string]process.Handler{
|
process.RegisterGroup("agent", map[string]process.Handler{
|
||||||
"write": ProcessWrite,
|
"write": ProcessWrite,
|
||||||
|
|
@ -201,11 +193,11 @@ func parseAssistantFilter(params map[string]interface{}) store.AssistantFilter {
|
||||||
case []interface{}:
|
case []interface{}:
|
||||||
filter.Select = []string{}
|
filter.Select = []string{}
|
||||||
for _, field := range v {
|
for _, field := range v {
|
||||||
switch field.(type) {
|
switch v := field.(type) {
|
||||||
case string:
|
case string:
|
||||||
filter.Select = append(filter.Select, field.(string))
|
filter.Select = append(filter.Select, v)
|
||||||
case interface{}:
|
case interface{}:
|
||||||
filter.Select = append(filter.Select, fmt.Sprintf("%v", field))
|
filter.Select = append(filter.Select, fmt.Sprintf("%v", v))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
11
agent/types/dsl.go
Normal file
11
agent/types/dsl.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package agent
|
package types
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"github.com/gin-gonic/gin"
|
"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
|
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
|
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
|
// AuthSetting *Auth `json:"auth,omitempty" yaml:"auth,omitempty"` // Authenticate Settings
|
||||||
// UploadSetting *Upload `json:"upload,omitempty" yaml:"upload,omitempty"` // Upload Settings
|
// UploadSetting *Upload `json:"upload,omitempty" yaml:"upload,omitempty"` // Upload Settings
|
||||||
// KnowledgeSetting *Knowledge `json:"knowledge,omitempty" yaml:"knowledge,omitempty"` // Knowledge base Settings
|
// KnowledgeSetting *Knowledge `json:"knowledge,omitempty" yaml:"knowledge,omitempty"` // Knowledge base Settings
|
||||||
|
|
@ -29,7 +31,7 @@ type DSL struct {
|
||||||
|
|
||||||
// Internal
|
// 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
|
Assistant assistant.API `json:"-" yaml:"-"` // The default assistant
|
||||||
Store store.Store `json:"-" yaml:"-"` // The store of the assistant
|
Store store.Store `json:"-" yaml:"-"` // The store of the assistant
|
||||||
Vision *vision.Vision `json:"-" yaml:"-"`
|
Vision *vision.Vision `json:"-" yaml:"-"`
|
||||||
|
|
@ -2,7 +2,6 @@ package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/yaoapp/yao/agent"
|
|
||||||
"github.com/yaoapp/yao/openapi/oauth/types"
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -11,19 +10,19 @@ import (
|
||||||
func Attach(group *gin.RouterGroup, oauth types.OAuth) {
|
func Attach(group *gin.RouterGroup, oauth types.OAuth) {
|
||||||
|
|
||||||
// Get the Agent instance
|
// Get the Agent instance
|
||||||
n := agent.GetAgent()
|
// n := agent.GetAgent()
|
||||||
|
|
||||||
// Apply OAuth guard to all routes
|
// Apply OAuth guard to all routes
|
||||||
group.Use(oauth.Guard)
|
group.Use(oauth.Guard)
|
||||||
|
|
||||||
// Assistant CRUD - Standard REST endpoints
|
// Assistant CRUD - Standard REST endpoints
|
||||||
group.GET("/assistants", ListAssistants) // GET /assistants - List assistants
|
group.GET("/assistants", ListAssistants) // GET /assistants - List assistants
|
||||||
group.POST("/assistants", CreateAssistant) // POST /assistants - Create assistant
|
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/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.GET("/assistants/:id", GetAssistant) // GET /assistants/:id - Get assistant details with permission verification
|
||||||
group.PUT("/assistants/:id", UpdateAssistant) // PUT /assistants/:id - Update assistant
|
group.PUT("/assistants/:id", UpdateAssistant) // PUT /assistants/:id - Update assistant
|
||||||
group.DELETE("/assistants/:id", n.HandleAssistantDelete) // DELETE /assistants/:id - Delete assistant
|
// group.DELETE("/assistants/:id", agent.HandleAssistantDelete) // DELETE /assistants/:id - Delete assistant
|
||||||
|
|
||||||
// Assistant Actions
|
// 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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -9,16 +9,45 @@ import (
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
|
"github.com/yaoapp/yao/agent"
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/openapi/response"
|
"github.com/yaoapp/yao/openapi/response"
|
||||||
)
|
)
|
||||||
|
|
||||||
// GinCreateCompletions handles POST /chat/:assistant_id/completions - Create a chat completion
|
// GinCreateCompletions handles POST /chat/:assistant_id/completions - Create a chat completion
|
||||||
func GinCreateCompletions(c *gin.Context) {
|
func GinCreateCompletions(c *gin.Context) {
|
||||||
// Print request information for debugging
|
|
||||||
fmt.Println("========== Chat Completions Request ==========")
|
agent := agent.GetAgent()
|
||||||
fmt.Printf("Method: %s\n", c.Request.Method)
|
cache, err := agent.GetCacheStore()
|
||||||
fmt.Printf("URL: %s\n", c.Request.URL.String())
|
if err != nil {
|
||||||
fmt.Printf("RemoteAddr: %s\n", c.Request.RemoteAddr)
|
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
|
// Print headers
|
||||||
fmt.Println("\n--- Headers ---")
|
fmt.Println("\n--- Headers ---")
|
||||||
|
|
@ -44,7 +73,7 @@ func GinCreateCompletions(c *gin.Context) {
|
||||||
|
|
||||||
// Print request body
|
// Print request body
|
||||||
fmt.Println("\n--- Request Body ---")
|
fmt.Println("\n--- Request Body ---")
|
||||||
body, err := io.ReadAll(c.Request.Body)
|
body, err = io.ReadAll(c.Request.Body)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Printf("Error reading body: %v\n", err)
|
fmt.Printf("Error reading body: %v\n", err)
|
||||||
} else {
|
} else {
|
||||||
|
|
|
||||||
|
|
@ -146,10 +146,10 @@ func (s *Service) getAccessTokenFromAPIKey(apiKey string) string {
|
||||||
log.Warn("Failed to make access token: %s", err.Error())
|
log.Warn("Failed to make access token: %s", err.Error())
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Println("========== Access Token From API Key ==========")
|
// fmt.Println("========== Access Token From API Key ==========")
|
||||||
fmt.Println("accessToken: ", accessToken)
|
// fmt.Println("accessToken: ", accessToken)
|
||||||
fmt.Println("extraClaims: ", extraClaims)
|
// fmt.Println("extraClaims: ", extraClaims)
|
||||||
fmt.Println("===============================================")
|
// fmt.Println("===============================================")
|
||||||
return accessToken
|
return accessToken
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@ import (
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/yaoapp/gou/api"
|
"github.com/yaoapp/gou/api"
|
||||||
"github.com/yaoapp/gou/server/http"
|
"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/config"
|
||||||
"github.com/yaoapp/yao/openapi"
|
"github.com/yaoapp/yao/openapi"
|
||||||
"github.com/yaoapp/yao/share"
|
"github.com/yaoapp/yao/share"
|
||||||
|
|
|
||||||
|
|
@ -223,6 +223,7 @@ var testSystemStores = map[string]string{
|
||||||
"__yao.oauth.client": "yao/stores/oauth/client.badger.yao",
|
"__yao.oauth.client": "yao/stores/oauth/client.badger.yao",
|
||||||
"__yao.oauth.cache": "yao/stores/oauth/cache.lru.yao",
|
"__yao.oauth.cache": "yao/stores/oauth/cache.lru.yao",
|
||||||
"__yao.agent.memory": "yao/stores/agent/memory.badger.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.store": "yao/stores/kb/store.badger.yao",
|
||||||
"__yao.kb.cache": "yao/stores/kb/cache.lru.yao",
|
"__yao.kb.cache": "yao/stores/kb/cache.lru.yao",
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -18,7 +18,7 @@ import (
|
||||||
"github.com/yaoapp/gou/session"
|
"github.com/yaoapp/gou/session"
|
||||||
"github.com/yaoapp/kun/exception"
|
"github.com/yaoapp/kun/exception"
|
||||||
"github.com/yaoapp/kun/log"
|
"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/agent/assistant"
|
||||||
"github.com/yaoapp/yao/config"
|
"github.com/yaoapp/yao/config"
|
||||||
"github.com/yaoapp/yao/data"
|
"github.com/yaoapp/yao/data"
|
||||||
|
|
@ -47,7 +47,7 @@ import (
|
||||||
|
|
||||||
// Setting the application setting
|
// Setting the application setting
|
||||||
var Setting *DSL
|
var Setting *DSL
|
||||||
var regExcp = regexp.MustCompile("^Exception\\|([0-9]+):(.+)$")
|
var regExcp = regexp.MustCompile(`^Exception\|([0-9]+):(.+)$`)
|
||||||
|
|
||||||
// LoadAndExport load app
|
// LoadAndExport load app
|
||||||
func LoadAndExport(cfg config.Config) error {
|
func LoadAndExport(cfg config.Config) error {
|
||||||
|
|
@ -258,7 +258,7 @@ func processService(process *process.Process) interface{} {
|
||||||
process.ValidateArgNums(2)
|
process.ValidateArgNums(2)
|
||||||
service := fmt.Sprintf("__yao_service.%s", process.ArgsString(0))
|
service := fmt.Sprintf("__yao_service.%s", process.ArgsString(0))
|
||||||
payload := process.ArgsMap(1)
|
payload := process.ArgsMap(1)
|
||||||
if payload == nil || len(payload) == 0 {
|
if len(payload) == 0 {
|
||||||
exception.New("content is required", 400).Throw()
|
exception.New("content is required", 400).Throw()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -508,7 +508,7 @@ func processXgen(process *process.Process) interface{} {
|
||||||
"layout": layout,
|
"layout": layout,
|
||||||
}
|
}
|
||||||
|
|
||||||
if admin.ThirdPartyLogin != nil && len(admin.ThirdPartyLogin) > 0 {
|
if len(admin.ThirdPartyLogin) > 0 {
|
||||||
xgenLogin["admin"]["thirdPartyLogin"] = admin.ThirdPartyLogin
|
xgenLogin["admin"]["thirdPartyLogin"] = admin.ThirdPartyLogin
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -544,7 +544,7 @@ func processXgen(process *process.Process) interface{} {
|
||||||
"layout": layout,
|
"layout": layout,
|
||||||
}
|
}
|
||||||
|
|
||||||
if user.ThirdPartyLogin != nil && len(user.ThirdPartyLogin) > 0 {
|
if len(user.ThirdPartyLogin) > 0 {
|
||||||
xgenLogin["user"]["thirdPartyLogin"] = user.ThirdPartyLogin
|
xgenLogin["user"]["thirdPartyLogin"] = user.ThirdPartyLogin
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -715,8 +715,9 @@ func (dsl *DSL) replaceAdminRoot() error {
|
||||||
// icons
|
// icons
|
||||||
func (dsl *DSL) icons(cfg config.Config) {
|
func (dsl *DSL) icons(cfg config.Config) {
|
||||||
|
|
||||||
dsl.Favicon = fmt.Sprintf("/api/__yao/app/icons/app.ico")
|
dsl.Favicon = "/api/__yao/app/icons/app.ico"
|
||||||
dsl.Logo = fmt.Sprintf("/api/__yao/app/icons/app.png")
|
dsl.Logo = "/api/__yao/app/icons/app.png"
|
||||||
|
log.Trace("CFG %v", cfg.Root)
|
||||||
|
|
||||||
// favicon := filepath.Join(cfg.Root, "icons", "app.ico")
|
// favicon := filepath.Join(cfg.Root, "icons", "app.ico")
|
||||||
// if _, err := os.Stat(favicon); err == nil {
|
// 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 {
|
for _, value := range values {
|
||||||
permissions[fmt.Sprintf("%v", value)] = true
|
permissions[fmt.Sprintf("%v", value)] = true
|
||||||
}
|
}
|
||||||
break
|
|
||||||
|
|
||||||
case []string:
|
case []string:
|
||||||
for _, value := range values {
|
for _, value := range values {
|
||||||
permissions[value] = true
|
permissions[value] = true
|
||||||
}
|
}
|
||||||
break
|
|
||||||
|
|
||||||
case map[string]interface{}:
|
case map[string]interface{}:
|
||||||
for key := range values {
|
for key := range values {
|
||||||
permissions[key] = true
|
permissions[key] = true
|
||||||
}
|
}
|
||||||
break
|
|
||||||
|
|
||||||
case map[string]bool:
|
case map[string]bool:
|
||||||
permissions = values
|
permissions = values
|
||||||
break
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return permissions
|
return permissions
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue