Merge pull request #1292 from trheyi/main

Refactor completion request handling in context tests
This commit is contained in:
Max 2025-11-11 16:24:01 +08:00 committed by GitHub
commit 858ee0368b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1803 additions and 474 deletions

View file

@ -1,186 +1,215 @@
package context package context
import ( import (
"bytes"
"encoding/json"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/yaoapp/gou/store"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test"
) )
func TestNewOpenAPI(t *testing.T) { func TestGetCompletionRequest(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
cache, err := store.Get("__yao.agent.cache")
if err != nil {
t.Fatalf("Failed to get cache: %v", err)
}
tests := []struct { tests := []struct {
name string name string
requestBody map[string]interface{}
queryParams map[string]string queryParams map[string]string
routeParams map[string]string
headers map[string]string headers map[string]string
expectedChatID string expectedModel string
expectedAssistant string expectedMsgCount int
expectedTemp *float64
expectedStream *bool
expectedLocale string expectedLocale string
expectedTheme string expectedTheme string
expectedClientType string
expectedReferer string expectedReferer string
expectedAccept Accept expectedAccept Accept
expectedAssistantID string
expectError bool
}{ }{
{ {
name: "Parse all query parameters", name: "Complete request from body with metadata",
queryParams: map[string]string{ requestBody: map[string]interface{}{
"assistant_id": "ast456", "model": "gpt-4-yao_assistant123",
"chat_id": "chat123", "messages": []map[string]interface{}{
"locale": "zh-CN", {"role": "user", "content": "Hello"},
"theme": "Dark",
"referer": RefererProcess,
"accept": string(AcceptStandard),
}, },
routeParams: map[string]string{ "temperature": 0.7,
"assistant_id": "route-ast", "stream": true,
"metadata": map[string]string{
"locale": "zh-cn",
"theme": "dark",
"referer": "process",
"accept": "cui-web",
"chat_id": "chat-from-metadata",
}, },
headers: map[string]string{
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
}, },
expectedChatID: "chat123", expectedModel: "gpt-4-yao_assistant123",
expectedAssistant: "ast456", expectedMsgCount: 1,
expectedTemp: floatPtr(0.7),
expectedStream: boolPtr(true),
expectedLocale: "zh-cn", expectedLocale: "zh-cn",
expectedTheme: "dark", expectedTheme: "dark",
expectedClientType: "macos",
expectedReferer: RefererProcess, expectedReferer: RefererProcess,
expectedAccept: AcceptStandard, expectedAccept: AcceptWebCUI,
expectedAssistantID: "assistant123",
expectError: false,
}, },
{ {
name: "Default values with no parameters", name: "Query params override payload metadata",
queryParams: map[string]string{}, requestBody: map[string]interface{}{
headers: map[string]string{ "model": "gpt-4-yao_test456",
"User-Agent": "Mozilla/5.0", "messages": []map[string]interface{}{
{"role": "user", "content": "Test"},
}, },
expectedChatID: "", "metadata": map[string]string{
expectedAssistant: "", "locale": "en-us",
expectedLocale: "", "theme": "light",
expectedTheme: "", },
expectedClientType: "web", },
queryParams: map[string]string{
"locale": "fr-FR",
"theme": "auto",
},
expectedModel: "gpt-4-yao_test456",
expectedMsgCount: 1,
expectedLocale: "fr-fr",
expectedTheme: "auto",
expectedReferer: RefererAPI, expectedReferer: RefererAPI,
expectedAccept: AcceptWebCUI, expectedAccept: AcceptWebCUI,
expectedAssistantID: "test456",
expectError: false,
}, },
{ {
name: "Android client type detection", name: "Headers override payload metadata",
queryParams: map[string]string{}, requestBody: map[string]interface{}{
"model": "gpt-3.5-turbo-yao_header789",
"messages": []map[string]interface{}{
{"role": "user", "content": "Test"},
},
"metadata": map[string]string{
"referer": "process",
"accept": "cui-web",
},
},
headers: map[string]string{ headers: map[string]string{
"User-Agent": "Mozilla/5.0 (Linux; Android 10)", "X-Yao-Referer": "mcp",
"X-Yao-Accept": "cui-desktop",
}, },
expectedClientType: "android", expectedModel: "gpt-3.5-turbo-yao_header789",
expectedReferer: RefererAPI, expectedMsgCount: 1,
expectedAccept: AccepNativeCUI, expectedLocale: "",
}, expectedTheme: "",
{
name: "iOS client type detection",
queryParams: map[string]string{},
headers: map[string]string{
"User-Agent": "Mozilla/5.0 (iPhone; CPU iPhone OS 14_0)",
},
expectedClientType: "ios",
expectedReferer: RefererAPI,
expectedAccept: AccepNativeCUI,
},
{
name: "Windows desktop client type detection",
queryParams: map[string]string{},
headers: map[string]string{
"User-Agent": "Mozilla/5.0 (Windows NT 10.0)",
},
expectedClientType: "windows",
expectedReferer: RefererAPI,
expectedAccept: AcceptDesktopCUI,
},
{
name: "Agent client type detection",
queryParams: map[string]string{},
headers: map[string]string{
"User-Agent": "Yao-Agent/1.0",
},
expectedClientType: "agent",
expectedReferer: RefererAPI,
expectedAccept: AcceptStandard,
},
{
name: "JSSDK client type detection",
queryParams: map[string]string{},
headers: map[string]string{
"User-Agent": "Yao-JSSDK/2.0",
},
expectedClientType: "jssdk",
expectedReferer: RefererAPI,
expectedAccept: AcceptStandard,
},
{
name: "Custom headers for referer and accept",
queryParams: map[string]string{},
headers: map[string]string{
"User-Agent": "Mozilla/5.0",
"X-Yao-Referer": RefererMCP,
"X-Yao-Accept": string(AcceptDesktopCUI),
},
expectedClientType: "web",
expectedReferer: RefererMCP, expectedReferer: RefererMCP,
expectedAccept: AcceptDesktopCUI, expectedAccept: AcceptDesktopCUI,
expectedAssistantID: "header789",
expectError: false,
}, },
{ {
name: "Query parameters override headers", name: "Minimal request without metadata",
queryParams: map[string]string{ requestBody: map[string]interface{}{
"referer": RefererJSSDK, "model": "gpt-4o-yao_minimal",
"accept": string(AcceptStandard), "messages": []map[string]interface{}{
{"role": "user", "content": "Hello"},
}, },
headers: map[string]string{
"User-Agent": "Mozilla/5.0",
"X-Yao-Referer": RefererMCP,
"X-Yao-Accept": string(AcceptDesktopCUI),
}, },
expectedClientType: "web", expectedModel: "gpt-4o-yao_minimal",
expectedReferer: RefererJSSDK, expectedMsgCount: 1,
expectedAccept: AcceptStandard, expectedLocale: "",
expectedTheme: "",
expectedReferer: RefererAPI,
expectedAccept: AcceptWebCUI,
expectedAssistantID: "minimal",
expectError: false,
},
{
name: "Missing model",
requestBody: map[string]interface{}{
"messages": []map[string]interface{}{
{"role": "user", "content": "Hello"},
},
},
expectError: true,
},
{
name: "Missing messages",
requestBody: map[string]interface{}{
"model": "gpt-4",
},
expectError: true,
}, },
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
// Create test server
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
// Build query string // Build request
req, _ := http.NewRequest("GET", "http://example.com/test", nil) bodyBytes, _ := json.Marshal(tt.requestBody)
req, _ := http.NewRequest("POST", "http://example.com/chat/completions", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
// Add query params
q := req.URL.Query() q := req.URL.Query()
for key, value := range tt.queryParams { for key, value := range tt.queryParams {
q.Add(key, value) q.Add(key, value)
} }
req.URL.RawQuery = q.Encode() req.URL.RawQuery = q.Encode()
// Set headers // Add headers
for key, value := range tt.headers { for key, value := range tt.headers {
req.Header.Set(key, value) req.Header.Set(key, value)
} }
c.Request = req c.Request = req
// Set route params // Call GetCompletionRequest
for key, value := range tt.routeParams { completionReq, ctx, err := GetCompletionRequest(c, cache)
c.Params = append(c.Params, gin.Param{Key: key, Value: value})
if tt.expectError {
assert.Error(t, err)
return
} }
// Call NewGin assert.NoError(t, err)
ctx := NewOpenAPI(c, nil) assert.NotNil(t, completionReq)
assert.NotNil(t, ctx)
// Assertions // Verify CompletionRequest
assert.Equal(t, tt.expectedChatID, ctx.ChatID, "ChatID mismatch") assert.Equal(t, tt.expectedModel, completionReq.Model)
assert.Equal(t, tt.expectedAssistant, ctx.AssistantID, "AssistantID mismatch") assert.Equal(t, tt.expectedMsgCount, len(completionReq.Messages))
assert.Equal(t, tt.expectedLocale, ctx.Locale, "Locale mismatch") if tt.expectedTemp != nil {
assert.Equal(t, tt.expectedTheme, ctx.Theme, "Theme mismatch") assert.NotNil(t, completionReq.Temperature)
assert.Equal(t, tt.expectedClientType, ctx.Client.Type, "Client.Type mismatch") assert.Equal(t, *tt.expectedTemp, *completionReq.Temperature)
assert.Equal(t, tt.expectedReferer, ctx.Referer, "Referer mismatch") }
assert.Equal(t, tt.expectedAccept, ctx.Accept, "Accept mismatch") if tt.expectedStream != nil {
assert.NotNil(t, ctx.Space, "Space should not be nil") assert.NotNil(t, completionReq.Stream)
// Client.UserAgent and Client.IP are set from headers/request, may be empty in test context assert.Equal(t, *tt.expectedStream, *completionReq.Stream)
}
// Verify Context
assert.Equal(t, tt.expectedLocale, ctx.Locale)
assert.Equal(t, tt.expectedTheme, ctx.Theme)
assert.Equal(t, tt.expectedReferer, ctx.Referer)
assert.Equal(t, tt.expectedAccept, ctx.Accept)
assert.Equal(t, tt.expectedAssistantID, ctx.AssistantID)
assert.NotNil(t, ctx.Space)
assert.NotNil(t, ctx.Cache)
}) })
} }
} }
@ -287,3 +316,12 @@ func TestValidateReferer(t *testing.T) {
}) })
} }
} }
// Helper functions
func floatPtr(f float64) *float64 {
return &f
}
func boolPtr(b bool) *bool {
return &b
}

View file

@ -2,6 +2,7 @@ package context
import ( import (
"bytes" "bytes"
"encoding/base64"
"encoding/json" "encoding/json"
"fmt" "fmt"
"io" "io"
@ -13,42 +14,58 @@ import (
"github.com/yaoapp/yao/openapi/oauth/authorized" "github.com/yaoapp/yao/openapi/oauth/authorized"
) )
// NewOpenAPI create a new context from openapi context // GetCompletionRequest parse completion request and create context from openapi request
func NewOpenAPI(c *gin.Context, cache store.Store) Context { // Returns: *CompletionRequest, *Context, error
func GetCompletionRequest(c *gin.Context, cache store.Store) (*CompletionRequest, *Context, error) {
// Get authorized information // Get authorized information
authInfo := authorized.GetInfo(c) authInfo := authorized.GetInfo(c)
// Extract assistant ID (route parameter takes priority, handled in GetAssistantID) // Parse completion request from payload or query first
assistantID, _ := GetAssistantID(c) completionReq, err := parseCompletionRequestData(c)
if err != nil {
return nil, nil, fmt.Errorf("failed to parse completion request: %w", err)
}
// Extract assistant ID using completionReq (can extract from model field)
assistantID, err := GetAssistantID(c, completionReq)
if err != nil {
return nil, nil, fmt.Errorf("failed to get assistant ID: %w", err)
}
// Extract chat ID (may generate from messages if not provided) // Extract chat ID (may generate from messages if not provided)
// GetChatID internally calls GetChatIDByMessages which auto-caches chatID, err := GetChatID(c, cache, completionReq)
chatID, _ := GetChatID(c, cache) if err != nil {
// If chat ID generation fails, it's not critical - allow empty chatID
chatID = ""
}
// Parse client information from User-Agent header // Parse client information from User-Agent header
userAgent := c.GetHeader("User-Agent") userAgent := c.GetHeader("User-Agent")
clientType := getClientType(userAgent) clientType := getClientType(userAgent)
clientIP := c.ClientIP() clientIP := c.ClientIP()
// Create context with extracted parameters // Set cache in context
ctx := Context{ ctx := &Context{
Context: c.Request.Context(), Context: c.Request.Context(),
Space: plan.NewMemorySharedSpace(), Space: plan.NewMemorySharedSpace(),
Cache: cache,
Authorized: authInfo, Authorized: authInfo,
ChatID: chatID, ChatID: chatID,
AssistantID: assistantID, AssistantID: assistantID,
Locale: GetLocale(c), Locale: GetLocale(c, completionReq),
Theme: GetTheme(c), Theme: GetTheme(c, completionReq),
Referer: GetReferer(c), Referer: GetReferer(c, completionReq),
Accept: GetAccept(c), Accept: GetAccept(c, completionReq),
Client: Client{ Client: Client{
Type: clientType, Type: clientType,
UserAgent: userAgent, UserAgent: userAgent,
IP: clientIP, IP: clientIP,
}, },
Route: GetRoute(c, completionReq),
Data: GetData(c, completionReq),
} }
return ctx return completionReq, ctx, nil
} }
// getClientType parses the client type from User-Agent header // getClientType parses the client type from User-Agent header
@ -80,45 +97,11 @@ func getClientType(userAgent string) string {
} }
} }
// 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: // GetAssistantID extracts assistant ID from request with priority:
// 1. Query parameter "assistant_id" // 1. Query parameter "assistant_id"
// 2. Header "X-Yao-Assistant" // 2. Header "X-Yao-Assistant"
// 3. Query parameter "model" - splits by "-" takes last field, extracts ID from "yao_xxx" prefix // 3. Extract from model field (from CompletionRequest or Query) - 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, req *CompletionRequest) (string, error) {
func GetAssistantID(c *gin.Context) (string, error) {
// Priority 1: Query parameter assistant_id // Priority 1: Query parameter assistant_id
if assistantID := c.Query("assistant_id"); assistantID != "" { if assistantID := c.Query("assistant_id"); assistantID != "" {
return assistantID, nil return assistantID, nil
@ -129,10 +112,12 @@ func GetAssistantID(c *gin.Context) (string, error) {
return assistantID, nil return assistantID, nil
} }
// Priority 3 & 4: Extract from model parameter (Query or Payload) // Priority 3: Extract from model field (from CompletionRequest or Query)
model := c.Query("model") model := ""
if model == "" { if req != nil && req.Model != "" {
model = getPayloadField(c, "model") model = req.Model
} else {
model = c.Query("model")
} }
if model != "" { if model != "" {
@ -156,9 +141,9 @@ func GetAssistantID(c *gin.Context) (string, error) {
// GetMessages extracts messages from the request // GetMessages extracts messages from the request
// Priority: // Priority:
// 1. Query parameter "messages" (JSON string) // 1. Query parameter "messages" (JSON string)
// 2. Request body "messages" field // 2. CompletionRequest.Messages (from payload)
func GetMessages(c *gin.Context) ([]Message, error) { func GetMessages(c *gin.Context, req *CompletionRequest) ([]Message, error) {
// Try query parameter first // Priority 1: Query parameter messages
if messagesJSON := c.Query("messages"); messagesJSON != "" { if messagesJSON := c.Query("messages"); messagesJSON != "" {
var messages []Message var messages []Message
if err := json.Unmarshal([]byte(messagesJSON), &messages); err == nil && len(messages) > 0 { if err := json.Unmarshal([]byte(messagesJSON), &messages); err == nil && len(messages) > 0 {
@ -166,45 +151,19 @@ func GetMessages(c *gin.Context) ([]Message, error) {
} }
} }
// Check if request body exists // Priority 2: From CompletionRequest (payload)
if c.Request.Body == nil { if req != nil && len(req.Messages) > 0 {
return req.Messages, nil
}
return nil, fmt.Errorf("messages field is required") 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: // GetLocale extracts locale from request with priority:
// 1. Query parameter "locale" // 1. Query parameter "locale"
// 2. Header "Accept-Language" // 2. Header "Accept-Language"
func GetLocale(c *gin.Context) string { // 3. CompletionRequest metadata "locale" (from payload)
func GetLocale(c *gin.Context, req *CompletionRequest) string {
// Priority 1: Query parameter // Priority 1: Query parameter
if locale := c.Query("locale"); locale != "" { if locale := c.Query("locale"); locale != "" {
return strings.ToLower(locale) return strings.ToLower(locale)
@ -222,13 +181,21 @@ func GetLocale(c *gin.Context) string {
} }
} }
// Priority 3: From CompletionRequest metadata
if req != nil && req.Metadata != nil {
if locale, ok := req.Metadata["locale"]; ok && locale != "" {
return strings.ToLower(locale)
}
}
return "" return ""
} }
// GetTheme extracts theme from request with priority: // GetTheme extracts theme from request with priority:
// 1. Query parameter "theme" // 1. Query parameter "theme"
// 2. Header "X-Yao-Theme" // 2. Header "X-Yao-Theme"
func GetTheme(c *gin.Context) string { // 3. CompletionRequest metadata "theme" (from payload)
func GetTheme(c *gin.Context, req *CompletionRequest) string {
// Priority 1: Query parameter // Priority 1: Query parameter
if theme := c.Query("theme"); theme != "" { if theme := c.Query("theme"); theme != "" {
return strings.ToLower(theme) return strings.ToLower(theme)
@ -239,14 +206,22 @@ func GetTheme(c *gin.Context) string {
return strings.ToLower(theme) return strings.ToLower(theme)
} }
// Priority 3: From CompletionRequest metadata
if req != nil && req.Metadata != nil {
if theme, ok := req.Metadata["theme"]; ok && theme != "" {
return strings.ToLower(theme)
}
}
return "" return ""
} }
// GetReferer extracts referer from request with priority: // GetReferer extracts referer from request with priority:
// 1. Query parameter "referer" // 1. Query parameter "referer"
// 2. Header "X-Yao-Referer" // 2. Header "X-Yao-Referer"
// 3. Default to "api" // 3. CompletionRequest metadata "referer" (from payload)
func GetReferer(c *gin.Context) string { // 4. Default to "api"
func GetReferer(c *gin.Context, req *CompletionRequest) string {
// Priority 1: Query parameter // Priority 1: Query parameter
if referer := c.Query("referer"); referer != "" { if referer := c.Query("referer"); referer != "" {
return validateReferer(referer) return validateReferer(referer)
@ -257,15 +232,23 @@ func GetReferer(c *gin.Context) string {
return validateReferer(referer) return validateReferer(referer)
} }
// Priority 3: Default // Priority 3: From CompletionRequest metadata
if req != nil && req.Metadata != nil {
if referer, ok := req.Metadata["referer"]; ok && referer != "" {
return validateReferer(referer)
}
}
// Priority 4: Default
return RefererAPI return RefererAPI
} }
// GetAccept extracts accept type from request with priority: // GetAccept extracts accept type from request with priority:
// 1. Query parameter "accept" // 1. Query parameter "accept"
// 2. Header "X-Yao-Accept" // 2. Header "X-Yao-Accept"
// 3. Parse from client type (User-Agent) // 3. CompletionRequest metadata "accept" (from payload)
func GetAccept(c *gin.Context) Accept { // 4. Parse from client type (User-Agent)
func GetAccept(c *gin.Context, req *CompletionRequest) Accept {
// Priority 1: Query parameter // Priority 1: Query parameter
if accept := c.Query("accept"); accept != "" { if accept := c.Query("accept"); accept != "" {
return validateAccept(accept) return validateAccept(accept)
@ -276,7 +259,14 @@ func GetAccept(c *gin.Context) Accept {
return validateAccept(accept) return validateAccept(accept)
} }
// Priority 3: Parse from User-Agent // Priority 3: From CompletionRequest metadata
if req != nil && req.Metadata != nil {
if accept, ok := req.Metadata["accept"]; ok && accept != "" {
return validateAccept(accept)
}
}
// Priority 4: Parse from User-Agent
userAgent := c.GetHeader("User-Agent") userAgent := c.GetHeader("User-Agent")
clientType := getClientType(userAgent) clientType := getClientType(userAgent)
return parseAccept(clientType) return parseAccept(clientType)
@ -286,8 +276,9 @@ func GetAccept(c *gin.Context) Accept {
// Priority: // Priority:
// 1. Query parameter "chat_id" // 1. Query parameter "chat_id"
// 2. Header "X-Yao-Chat" // 2. Header "X-Yao-Chat"
// 3. Generate from messages using GetChatIDByMessages // 3. CompletionRequest metadata "chat_id" (from payload)
func GetChatID(c *gin.Context, cache store.Store) (string, error) { // 4. Generate from messages using GetChatIDByMessages
func GetChatID(c *gin.Context, cache store.Store, req *CompletionRequest) (string, error) {
// Priority 1: Query parameter chat_id // Priority 1: Query parameter chat_id
if chatID := c.Query("chat_id"); chatID != "" { if chatID := c.Query("chat_id"); chatID != "" {
return chatID, nil return chatID, nil
@ -298,8 +289,15 @@ func GetChatID(c *gin.Context, cache store.Store) (string, error) {
return chatID, nil return chatID, nil
} }
// Priority 3: Generate from messages // Priority 3: From CompletionRequest metadata
messages, err := GetMessages(c) if req != nil && req.Metadata != nil {
if chatID, ok := req.Metadata["chat_id"]; ok && chatID != "" {
return chatID, nil
}
}
// Priority 4: Generate from messages
messages, err := GetMessages(c, req)
if err != nil { if err != nil {
return "", fmt.Errorf("failed to get messages for chat ID generation: %w", err) return "", fmt.Errorf("failed to get messages for chat ID generation: %w", err)
} }
@ -311,3 +309,170 @@ func GetChatID(c *gin.Context, cache store.Store) (string, error) {
return chatID, nil return chatID, nil
} }
// GetRoute extracts route from request with priority:
// 1. Query parameter "yao_route"
// 2. Header "X-Yao-Route"
// 3. CompletionRequest.Route (from payload)
func GetRoute(c *gin.Context, req *CompletionRequest) string {
// Priority 1: Query parameter
if route := c.Query("yao_route"); route != "" {
return route
}
// Priority 2: Header
if route := c.GetHeader("X-Yao-Route"); route != "" {
return route
}
// Priority 3: From CompletionRequest
if req != nil && req.Route != "" {
return req.Route
}
return ""
}
// GetData extracts data from request with priority:
// 1. Query parameter "yao_data" (JSON string)
// 2. Header "X-Yao-Data" (Base64 encoded JSON string)
// 3. CompletionRequest.Data (from payload)
func GetData(c *gin.Context, req *CompletionRequest) map[string]interface{} {
// Priority 1: Query parameter (JSON string)
if dataJSON := c.Query("yao_data"); dataJSON != "" {
var data map[string]interface{}
if err := json.Unmarshal([]byte(dataJSON), &data); err == nil {
return data
}
}
// Priority 2: Header (Base64 encoded JSON string)
if dataBase64 := c.GetHeader("X-Yao-Data"); dataBase64 != "" {
// Try to decode Base64
if decoded, err := base64.StdEncoding.DecodeString(dataBase64); err == nil {
var data map[string]interface{}
if err := json.Unmarshal(decoded, &data); err == nil {
return data
}
}
// Fallback: try to parse as plain JSON (for backward compatibility)
var data map[string]interface{}
if err := json.Unmarshal([]byte(dataBase64), &data); err == nil {
return data
}
}
// Priority 3: From CompletionRequest
if req != nil && req.Data != nil {
return req.Data
}
return nil
}
// parseCompletionRequestData extracts CompletionRequest from the request
// Data can be passed via:
// 1. Request body (JSON payload) - Priority
// 2. Query parameters
func parseCompletionRequestData(c *gin.Context) (*CompletionRequest, error) {
var req CompletionRequest
// Try to parse from request body first
if c.Request.Body != nil {
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 not empty, try to parse it
if len(body) > 0 {
if err := json.Unmarshal(body, &req); err != nil {
return nil, fmt.Errorf("failed to parse completion request from body: %w", err)
}
// If we got valid data from body, validate and return
if req.Model != "" && len(req.Messages) > 0 {
return &req, nil
}
}
}
// Fallback: Try to parse from query parameters
// Required fields
model := c.Query("model")
if model == "" {
return nil, fmt.Errorf("model field is required")
}
req.Model = model
// Messages (required, must be JSON string in query)
messagesJSON := c.Query("messages")
if messagesJSON == "" {
return nil, fmt.Errorf("messages field is required")
}
var messages []Message
if err := json.Unmarshal([]byte(messagesJSON), &messages); err != nil {
return nil, fmt.Errorf("failed to parse messages from query: %w", err)
}
if len(messages) == 0 {
return nil, fmt.Errorf("messages field must not be empty")
}
req.Messages = messages
// Optional fields from query
if tempStr := c.Query("temperature"); tempStr != "" {
var temp float64
if _, err := fmt.Sscanf(tempStr, "%f", &temp); err == nil {
req.Temperature = &temp
}
}
if maxTokensStr := c.Query("max_tokens"); maxTokensStr != "" {
var maxTokens int
if _, err := fmt.Sscanf(maxTokensStr, "%d", &maxTokens); err == nil {
req.MaxTokens = &maxTokens
}
}
if maxCompletionTokensStr := c.Query("max_completion_tokens"); maxCompletionTokensStr != "" {
var maxCompletionTokens int
if _, err := fmt.Sscanf(maxCompletionTokensStr, "%d", &maxCompletionTokens); err == nil {
req.MaxCompletionTokens = &maxCompletionTokens
}
}
if streamStr := c.Query("stream"); streamStr != "" {
stream := streamStr == "true" || streamStr == "1"
req.Stream = &stream
}
// Audio config from query (JSON string)
if audioJSON := c.Query("audio"); audioJSON != "" {
var audio AudioConfig
if err := json.Unmarshal([]byte(audioJSON), &audio); err == nil {
req.Audio = &audio
}
}
// Stream options from query (JSON string)
if streamOptionsJSON := c.Query("stream_options"); streamOptionsJSON != "" {
var streamOptions StreamOptions
if err := json.Unmarshal([]byte(streamOptionsJSON), &streamOptions); err == nil {
req.StreamOptions = &streamOptions
}
}
// Metadata from query (JSON string)
if metadataJSON := c.Query("metadata"); metadataJSON != "" {
var metadata map[string]string
if err := json.Unmarshal([]byte(metadataJSON), &metadata); err == nil {
req.Metadata = metadata
}
}
return &req, nil
}

View file

@ -42,7 +42,10 @@ func TestGetMessages_FromBody(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
result, err := GetMessages(c) // Parse request first
completionReq, _ := parseCompletionRequestData(c)
result, err := GetMessages(c, completionReq)
if err != nil { if err != nil {
t.Fatalf("Failed to get messages: %v", err) t.Fatalf("Failed to get messages: %v", err)
} }
@ -80,7 +83,7 @@ func TestGetMessages_FromQuery(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
result, err := GetMessages(c) result, err := GetMessages(c, nil)
if err != nil { if err != nil {
t.Fatalf("Failed to get messages: %v", err) t.Fatalf("Failed to get messages: %v", err)
} }
@ -109,7 +112,9 @@ func TestGetMessages_EmptyMessages(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
_, err := GetMessages(c) completionReq, _ := parseCompletionRequestData(c)
_, err := GetMessages(c, completionReq)
if err == nil { if err == nil {
t.Error("Expected error for empty messages") t.Error("Expected error for empty messages")
} }
@ -133,7 +138,7 @@ func TestGetChatID_FromQuery(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
chatID, err := GetChatID(c, cache) chatID, err := GetChatID(c, cache, nil)
if err != nil { if err != nil {
t.Fatalf("Failed to get chat ID: %v", err) t.Fatalf("Failed to get chat ID: %v", err)
} }
@ -162,7 +167,50 @@ func TestGetChatID_FromHeader(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
chatID, err := GetChatID(c, cache) chatID, err := GetChatID(c, cache, nil)
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_FromMetadata(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 := "metadata-chat-789"
requestBody := map[string]interface{}{
"model": "gpt-4",
"messages": []map[string]interface{}{
{"role": "user", "content": "Test"},
},
"metadata": map[string]string{
"chat_id": expectedChatID,
},
}
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
completionReq, _ := parseCompletionRequestData(c)
chatID, err := GetChatID(c, cache, completionReq)
if err != nil { if err != nil {
t.Fatalf("Failed to get chat ID: %v", err) t.Fatalf("Failed to get chat ID: %v", err)
} }
@ -193,6 +241,7 @@ func TestGetChatID_FromMessages(t *testing.T) {
} }
requestBody1 := map[string]interface{}{ requestBody1 := map[string]interface{}{
"model": "gpt-4",
"messages": messages1, "messages": messages1,
} }
@ -204,7 +253,9 @@ func TestGetChatID_FromMessages(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
chatID1, err := GetChatID(c, cache) completionReq1, _ := parseCompletionRequestData(c)
chatID1, err := GetChatID(c, cache, completionReq1)
if err != nil { if err != nil {
t.Fatalf("Failed to get chat ID: %v", err) t.Fatalf("Failed to get chat ID: %v", err)
} }
@ -226,6 +277,7 @@ func TestGetChatID_FromMessages(t *testing.T) {
} }
requestBody2 := map[string]interface{}{ requestBody2 := map[string]interface{}{
"model": "gpt-4",
"messages": messages2, "messages": messages2,
} }
@ -237,7 +289,9 @@ func TestGetChatID_FromMessages(t *testing.T) {
c2, _ := gin.CreateTestContext(w2) c2, _ := gin.CreateTestContext(w2)
c2.Request = req2 c2.Request = req2
chatID2, err := GetChatID(c2, cache) completionReq2, _ := parseCompletionRequestData(c2)
chatID2, err := GetChatID(c2, cache, completionReq2)
if err != nil { if err != nil {
t.Fatalf("Failed to get chat ID second time: %v", err) t.Fatalf("Failed to get chat ID second time: %v", err)
} }
@ -261,6 +315,7 @@ func TestGetChatID_Priority(t *testing.T) {
queryChatID := "query-chat-id" queryChatID := "query-chat-id"
headerChatID := "header-chat-id" headerChatID := "header-chat-id"
metadataChatID := "metadata-chat-id"
messages := []Message{ messages := []Message{
{ {
@ -270,12 +325,16 @@ func TestGetChatID_Priority(t *testing.T) {
} }
requestBody := map[string]interface{}{ requestBody := map[string]interface{}{
"model": "gpt-4",
"messages": messages, "messages": messages,
"metadata": map[string]string{
"chat_id": metadataChatID,
},
} }
bodyBytes, _ := json.Marshal(requestBody) bodyBytes, _ := json.Marshal(requestBody)
// Test priority: query > header > messages // Test priority: query > header > metadata > messages
req := httptest.NewRequest("POST", "/chat/completions?chat_id="+queryChatID, bytes.NewBuffer(bodyBytes)) req := httptest.NewRequest("POST", "/chat/completions?chat_id="+queryChatID, bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Yao-Chat", headerChatID) req.Header.Set("X-Yao-Chat", headerChatID)
@ -283,7 +342,9 @@ func TestGetChatID_Priority(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
chatID, err := GetChatID(c, cache) completionReq, _ := parseCompletionRequestData(c)
chatID, err := GetChatID(c, cache, completionReq)
if err != nil { if err != nil {
t.Fatalf("Failed to get chat ID: %v", err) t.Fatalf("Failed to get chat ID: %v", err)
} }
@ -301,7 +362,7 @@ func TestGetLocale_FromQuery(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
locale := GetLocale(c) locale := GetLocale(c, nil)
if locale != "zh-cn" { if locale != "zh-cn" {
t.Errorf("Expected locale 'zh-cn', got '%s'", locale) t.Errorf("Expected locale 'zh-cn', got '%s'", locale)
} }
@ -316,12 +377,32 @@ func TestGetLocale_FromHeader(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
locale := GetLocale(c) locale := GetLocale(c, nil)
if locale != "en-us" { if locale != "en-us" {
t.Errorf("Expected locale 'en-us', got '%s'", locale) t.Errorf("Expected locale 'en-us', got '%s'", locale)
} }
} }
func TestGetLocale_FromMetadata(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
completionReq := &CompletionRequest{
Metadata: map[string]string{
"locale": "ja-JP",
},
}
locale := GetLocale(c, completionReq)
if locale != "ja-jp" {
t.Errorf("Expected locale 'ja-jp' from metadata, got '%s'", locale)
}
}
func TestGetLocale_Priority(t *testing.T) { func TestGetLocale_Priority(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
@ -331,26 +412,18 @@ func TestGetLocale_Priority(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
locale := GetLocale(c) completionReq := &CompletionRequest{
Metadata: map[string]string{
"locale": "de-DE",
},
}
locale := GetLocale(c, completionReq)
if locale != "fr-fr" { if locale != "fr-fr" {
t.Errorf("Expected query parameter to take priority, got '%s'", locale) 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) { func TestGetTheme_FromQuery(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
@ -359,7 +432,7 @@ func TestGetTheme_FromQuery(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
theme := GetTheme(c) theme := GetTheme(c, nil)
if theme != "dark" { if theme != "dark" {
t.Errorf("Expected theme 'dark', got '%s'", theme) t.Errorf("Expected theme 'dark', got '%s'", theme)
} }
@ -374,213 +447,311 @@ func TestGetTheme_FromHeader(t *testing.T) {
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
theme := GetTheme(c) theme := GetTheme(c, nil)
if theme != "light" { if theme != "light" {
t.Errorf("Expected theme 'light', got '%s'", theme) t.Errorf("Expected theme 'light', got '%s'", theme)
} }
} }
func TestGetTheme_Priority(t *testing.T) { func TestGetTheme_FromMetadata(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions?theme=auto", nil) req := httptest.NewRequest("POST", "/chat/completions", nil)
req.Header.Set("X-Yao-Theme", "dark")
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
theme := GetTheme(c) completionReq := &CompletionRequest{
Metadata: map[string]string{
"theme": "auto",
},
}
theme := GetTheme(c, completionReq)
if theme != "auto" { if theme != "auto" {
t.Errorf("Expected query parameter to take priority, got '%s'", theme) t.Errorf("Expected theme 'auto' from metadata, got '%s'", theme)
} }
} }
func TestGetTheme_Empty(t *testing.T) { func TestGetReferer_FromMetadata(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions", nil) req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
theme := GetTheme(c) completionReq := &CompletionRequest{
if theme != "" { Metadata: map[string]string{
t.Errorf("Expected empty theme, got '%s'", theme) "referer": "tool",
},
}
referer := GetReferer(c, completionReq)
if referer != RefererTool {
t.Errorf("Expected referer 'tool' from metadata, got '%s'", referer)
} }
} }
func TestGetReferer_FromQuery(t *testing.T) { func TestGetAccept_FromMetadata(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions?referer=jssdk", nil) req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
referer := GetReferer(c) completionReq := &CompletionRequest{
if referer != "jssdk" { Metadata: map[string]string{
t.Errorf("Expected referer 'jssdk', got '%s'", referer) "accept": "cui-native",
} },
} }
func TestGetReferer_FromHeader(t *testing.T) { accept := GetAccept(c, completionReq)
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 { if accept != AccepNativeCUI {
t.Errorf("Expected accept '%s', got '%s'", AccepNativeCUI, accept) t.Errorf("Expected accept 'cui-native' from metadata, got '%s'", accept)
} }
} }
func TestGetAccept_FromUserAgent_Web(t *testing.T) { func TestGetAssistantID_FromModel(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
completionReq := &CompletionRequest{
Model: "gpt-4-turbo-yao_myassistant",
}
assistantID, err := GetAssistantID(c, completionReq)
if err != nil {
t.Fatalf("Failed to get assistant ID: %v", err)
}
if assistantID != "myassistant" {
t.Errorf("Expected assistant ID 'myassistant', got '%s'", assistantID)
}
}
func TestGetAssistantID_Priority(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions?assistant_id=from_query", nil)
req.Header.Set("X-Yao-Assistant", "from_header")
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
completionReq := &CompletionRequest{
Model: "gpt-4-yao_from_model",
}
assistantID, err := GetAssistantID(c, completionReq)
if err != nil {
t.Fatalf("Failed to get assistant ID: %v", err)
}
if assistantID != "from_query" {
t.Errorf("Expected query parameter to take priority, got '%s'", assistantID)
}
}
func TestGetRoute_FromQuery(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions?yao_route=/dashboard/home", nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
route := GetRoute(c, nil)
if route != "/dashboard/home" {
t.Errorf("Expected route '/dashboard/home', got '%s'", route)
}
}
func TestGetRoute_FromHeader(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions", nil) req := httptest.NewRequest("GET", "/chat/completions", nil)
req.Header.Set("User-Agent", "Mozilla/5.0") req.Header.Set("X-Yao-Route", "/settings/profile")
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
accept := GetAccept(c) route := GetRoute(c, nil)
if accept != AcceptWebCUI { if route != "/settings/profile" {
t.Errorf("Expected accept '%s' for web user agent, got '%s'", AcceptWebCUI, accept) t.Errorf("Expected route '/settings/profile', got '%s'", route)
} }
} }
func TestGetAccept_FromUserAgent_Android(t *testing.T) { func TestGetRoute_FromPayload(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
completionReq := &CompletionRequest{
Route: "/admin/users",
}
route := GetRoute(c, completionReq)
if route != "/admin/users" {
t.Errorf("Expected route '/admin/users' from payload, got '%s'", route)
}
}
func TestGetRoute_Priority(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions?yao_route=/from/query", nil)
req.Header.Set("X-Yao-Route", "/from/header")
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
completionReq := &CompletionRequest{
Route: "/from/payload",
}
route := GetRoute(c, completionReq)
if route != "/from/query" {
t.Errorf("Expected query parameter to take priority, got '%s'", route)
}
}
func TestGetData_FromQuery(t *testing.T) {
gin.SetMode(gin.TestMode)
data := map[string]interface{}{
"key1": "value1",
"key2": float64(123),
}
dataJSON, _ := json.Marshal(data)
req := httptest.NewRequest("GET", "/chat/completions?yao_data="+string(dataJSON), nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
result := GetData(c, nil)
if result == nil {
t.Fatal("Expected data to be returned")
}
if result["key1"] != "value1" {
t.Errorf("Expected key1='value1', got '%v'", result["key1"])
}
if result["key2"] != float64(123) {
t.Errorf("Expected key2=123, got '%v'", result["key2"])
}
}
func TestGetData_FromHeader_Base64(t *testing.T) {
gin.SetMode(gin.TestMode)
dataBase64 := "eyJ1c2VyX2lkIjo0NTYsImFjdGlvbiI6ImNyZWF0ZSJ9" // base64 of {"user_id":456,"action":"create"}
req := httptest.NewRequest("GET", "/chat/completions", nil)
req.Header.Set("X-Yao-Data", dataBase64)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
result := GetData(c, nil)
if result == nil {
t.Fatal("Expected data to be returned")
}
if result["action"] != "create" {
t.Errorf("Expected action='create', got '%v'", result["action"])
}
if result["user_id"] != float64(456) {
t.Errorf("Expected user_id=456, got '%v'", result["user_id"])
}
}
func TestGetData_FromPayload(t *testing.T) {
gin.SetMode(gin.TestMode)
req := httptest.NewRequest("POST", "/chat/completions", nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
data := map[string]interface{}{
"page": float64(1),
"limit": float64(10),
}
completionReq := &CompletionRequest{
Data: data,
}
result := GetData(c, completionReq)
if result == nil {
t.Fatal("Expected data to be returned")
}
if result["page"] != float64(1) {
t.Errorf("Expected page=1, got '%v'", result["page"])
}
if result["limit"] != float64(10) {
t.Errorf("Expected limit=10, got '%v'", result["limit"])
}
}
func TestGetData_Priority(t *testing.T) {
gin.SetMode(gin.TestMode)
queryData := map[string]interface{}{
"source": "query",
}
queryDataJSON, _ := json.Marshal(queryData)
headerDataBase64 := "eyJzb3VyY2UiOiJoZWFkZXIifQ==" // base64 of {"source":"header"}
req := httptest.NewRequest("GET", "/chat/completions?yao_data="+string(queryDataJSON), nil)
req.Header.Set("X-Yao-Data", headerDataBase64)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
payloadData := map[string]interface{}{
"source": "payload",
}
completionReq := &CompletionRequest{
Data: payloadData,
}
result := GetData(c, completionReq)
if result == nil {
t.Fatal("Expected data to be returned")
}
if result["source"] != "query" {
t.Errorf("Expected query parameter to take priority, got '%v'", result["source"])
}
}
func TestGetData_EmptyData(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
req := httptest.NewRequest("GET", "/chat/completions", nil) req := httptest.NewRequest("GET", "/chat/completions", nil)
req.Header.Set("User-Agent", "Android App")
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
accept := GetAccept(c) result := GetData(c, nil)
if accept != AccepNativeCUI { if result != nil {
t.Errorf("Expected accept '%s' for Android, got '%s'", AccepNativeCUI, accept) t.Errorf("Expected nil data, got '%v'", result)
}
}
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)
} }
} }

View file

@ -118,8 +118,8 @@ type Context struct {
Accept Accept `json:"accept,omitempty"` // Response format: standard, cui-web, cui-native, cui-desktop Accept Accept `json:"accept,omitempty"` // Response format: standard, cui-web, cui-native, cui-desktop
// CUI Context information // CUI Context information
Route string `json:"route,omitempty"` // The route of the request, it will be used to identify the route of the request Route string `json:"yao_route,omitempty"` // The route of the request, it will be used to identify the route of the request
Data map[string]interface{} `json:"data,omitempty"` // The data of the request, it will be used to pass data to the page Data map[string]interface{} `json:"yao_data,omitempty"` // The data of the request, it will be used to pass data to the page
Silent bool `json:"silent,omitempty"` // Silent mode (Deprecated, use Referer instead) Silent bool `json:"silent,omitempty"` // Silent mode (Deprecated, use Referer instead)
} }
@ -216,3 +216,43 @@ type Function struct {
Name string `json:"name"` // Required: name of the function to call 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 Arguments string `json:"arguments,omitempty"` // Optional: arguments to pass to the function, as a JSON string
} }
// Completion Request Structure ( OpenAI Chat Completion Request, https://platform.openai.com/docs/api-reference/chat/create )
// ===============================
// CompletionRequest represents a chat completion request compatible with OpenAI's API
type CompletionRequest struct {
// Required fields
Model string `json:"model"` // Required: ID of the model to use
Messages []Message `json:"messages"` // Required: list of messages comprising the conversation so far
// Audio configuration (for models that support audio output)
Audio *AudioConfig `json:"audio,omitempty"` // Optional: audio output configuration
// Generation parameters
Temperature *float64 `json:"temperature,omitempty"` // Optional: sampling temperature (0-2), defaults to 1
MaxTokens *int `json:"max_tokens,omitempty"` // Optional: maximum number of tokens to generate (deprecated, use max_completion_tokens)
MaxCompletionTokens *int `json:"max_completion_tokens,omitempty"` // Optional: maximum number of tokens that can be generated in the completion
// Streaming configuration
Stream *bool `json:"stream,omitempty"` // Optional: if true, stream partial message deltas
StreamOptions *StreamOptions `json:"stream_options,omitempty"` // Optional: options for streaming response
// Request metadata
Metadata map[string]string `json:"metadata,omitempty"` // Optional: developer-defined tags and values for tracking requests
// CUI Context information
Route string `json:"yao_route,omitempty"` // Optional: route of the request for CUI context
Data map[string]interface{} `json:"yao_data,omitempty"` // Optional: data to pass to the page for CUI context
}
// AudioConfig represents the audio output configuration for models that support audio
type AudioConfig struct {
Voice string `json:"voice"` // Required: voice to use for audio output (e.g., "alloy", "echo", "fable", "onyx", "nova", "shimmer")
Format string `json:"format"` // Required: audio output format (e.g., "wav", "mp3", "flac", "opus", "pcm16")
}
// StreamOptions represents options for streaming responses
type StreamOptions struct {
IncludeUsage bool `json:"include_usage,omitempty"` // If true, include usage statistics in the final chunk
}

View file

@ -2,10 +2,13 @@ package types
import ( import (
"fmt" "fmt"
"strings"
"time" "time"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
"github.com/spf13/cast" "github.com/spf13/cast"
"github.com/yaoapp/gou/connector"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/agent/i18n" "github.com/yaoapp/yao/agent/i18n"
) )
@ -374,3 +377,82 @@ func getBoolValue(data map[string]interface{}, key string) bool {
} }
return false return false
} }
// ModelID generates an OpenAI-compatible model ID from assistant
// Format: [prefix-]assistantName-model-yao_assistantID
// prefix is optional, if provided, it will be prepended to the model ID
func (assistant AssistantModel) ModelID(prefix ...string) string {
// Clean assistant name (remove spaces and special characters)
assistantName := strings.ReplaceAll(assistant.Name, " ", "-")
assistantName = strings.ToLower(assistantName)
// Get connector name from assistant
connectorName := assistant.Connector
if connectorName == "" {
log.Error("Assistant %s has no connector configured", assistant.ID)
modelID := assistantName + "-unknown-yao_" + assistant.ID
if len(prefix) > 0 && prefix[0] != "" {
return prefix[0] + modelID
}
return modelID
}
// Get model name
modelName := ""
// First, try to get custom model from Options
if assistant.Options != nil {
if m, ok := assistant.Options["model"].(string); ok && m != "" {
modelName = m
}
}
// If no custom model in options, try to get from connector configuration
if modelName == "" {
conn, err := connector.Select(connectorName)
if err != nil {
log.Error("Failed to select connector %s for assistant %s: %v", connectorName, assistant.ID, err)
modelID := assistantName + "-unknown-yao_" + assistant.ID
if len(prefix) > 0 && prefix[0] != "" {
return prefix[0] + modelID
}
return modelID
}
// Get model from connector settings
settings := conn.Setting()
if settings != nil {
if m, ok := settings["model"].(string); ok && m != "" {
modelName = m
}
}
if modelName == "" {
log.Error("Connector %s has no model configured for assistant %s", connectorName, assistant.ID)
modelID := assistantName + "-unknown-yao_" + assistant.ID
if len(prefix) > 0 && prefix[0] != "" {
return prefix[0] + modelID
}
return modelID
}
}
// Format: [prefix-]assistantName-model-yao_assistantID
modelID := assistantName + "-" + modelName + "-yao_" + assistant.ID
if len(prefix) > 0 && prefix[0] != "" {
return prefix[0] + modelID
}
return modelID
}
// ParseModelID extracts assistant ID from model ID
// Expected format: [prefix-]assistantName-model-yao_assistantID
// The function handles optional prefixes (e.g., "yao-agents-")
func ParseModelID(modelID string) string {
// Find the last occurrence of "yao_"
parts := strings.Split(modelID, "-yao_")
if len(parts) < 2 {
return ""
}
return parts[len(parts)-1]
}

View file

@ -867,3 +867,181 @@ func TestGetBoolValue(t *testing.T) {
} }
}) })
} }
// TestModelID tests the AssistantModel.ModelID method
func TestModelID(t *testing.T) {
t.Run("WithCustomModel", func(t *testing.T) {
assistant := AssistantModel{
ID: "test123",
Name: "Test Assistant",
Connector: "openai",
Options: map[string]interface{}{
"model": "gpt-4o",
},
}
result := assistant.ModelID()
expected := "test-assistant-gpt-4o-yao_test123"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithModelInOptions", func(t *testing.T) {
assistant := AssistantModel{
ID: "abc456",
Name: "My Bot",
Connector: "openai",
Options: map[string]interface{}{
"model": "gpt-3.5-turbo",
},
}
result := assistant.ModelID()
expected := "my-bot-gpt-3.5-turbo-yao_abc456"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithoutCustomModel", func(t *testing.T) {
assistant := AssistantModel{
ID: "xyz789",
Name: "Default Assistant",
Connector: "openai",
}
result := assistant.ModelID()
// When connector is not loaded in test, it should return unknown
expected := "default-assistant-unknown-yao_xyz789"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithoutConnector", func(t *testing.T) {
assistant := AssistantModel{
ID: "noconn",
Name: "No Connector",
}
result := assistant.ModelID()
expected := "no-connector-unknown-yao_noconn"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithSpacesInName", func(t *testing.T) {
assistant := AssistantModel{
ID: "spaces",
Name: "Test Bot With Spaces",
Connector: "anthropic",
Options: map[string]interface{}{
"model": "claude-3",
},
}
result := assistant.ModelID()
expected := "test-bot-with-spaces-claude-3-yao_spaces"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithUpperCaseName", func(t *testing.T) {
assistant := AssistantModel{
ID: "upper",
Name: "UPPERCASE-NAME",
Connector: "openai",
Options: map[string]interface{}{
"model": "GPT-4",
},
}
result := assistant.ModelID()
expected := "uppercase-name-GPT-4-yao_upper"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("WithEmptyOptions", func(t *testing.T) {
assistant := AssistantModel{
ID: "empty",
Name: "Empty Options",
Connector: "openai",
Options: map[string]interface{}{},
}
result := assistant.ModelID()
// When connector is not loaded in test, it should return unknown
expected := "empty-options-unknown-yao_empty"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
}
// TestParseModelID tests the ParseModelID function
func TestParseModelID(t *testing.T) {
t.Run("ValidModelID", func(t *testing.T) {
modelID := "test-assistant-gpt-4o-yao_test123"
result := ParseModelID(modelID)
expected := "test123"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("ValidModelIDWithMultipleDashes", func(t *testing.T) {
modelID := "my-test-bot-gpt-3.5-turbo-yao_abc456"
result := ParseModelID(modelID)
expected := "abc456"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("ValidModelIDWithHyphenInID", func(t *testing.T) {
modelID := "assistant-name-model-yao_id-with-dash"
result := ParseModelID(modelID)
expected := "id-with-dash"
if result != expected {
t.Errorf("Expected '%s', got '%s'", expected, result)
}
})
t.Run("InvalidModelIDNoYaoPrefix", func(t *testing.T) {
modelID := "test-assistant-gpt-4o-test123"
result := ParseModelID(modelID)
if result != "" {
t.Errorf("Expected empty string, got '%s'", result)
}
})
t.Run("InvalidModelIDEmpty", func(t *testing.T) {
modelID := ""
result := ParseModelID(modelID)
if result != "" {
t.Errorf("Expected empty string, got '%s'", result)
}
})
t.Run("InvalidModelIDOnlyYaoPrefix", func(t *testing.T) {
modelID := "yao_"
result := ParseModelID(modelID)
if result != "" {
t.Errorf("Expected empty string, got '%s'", result)
}
})
t.Run("RoundTrip", func(t *testing.T) {
assistant := AssistantModel{
ID: "roundtrip123",
Name: "Round Trip Test",
Connector: "openai",
Options: map[string]interface{}{
"model": "gpt-4",
},
}
modelID := assistant.ModelID()
extractedID := ParseModelID(modelID)
if extractedID != assistant.ID {
t.Errorf("Round trip failed: expected '%s', got '%s'", assistant.ID, extractedID)
}
})
}

View file

@ -58,8 +58,26 @@ func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
data["public"] = assistant.Public data["public"] = assistant.Public
data["mentionable"] = assistant.Mentionable data["mentionable"] = assistant.Mentionable
data["automated"] = assistant.Automated data["automated"] = assistant.Automated
data["created_at"] = assistant.CreatedAt
// Set timestamps
now := time.Now().UnixNano()
if exists {
// Update: set updated_at, keep created_at unchanged
if assistant.UpdatedAt == 0 {
data["updated_at"] = now
} else {
data["updated_at"] = assistant.UpdatedAt data["updated_at"] = assistant.UpdatedAt
}
// Don't modify created_at on update
} else {
// Create: set created_at, updated_at is null
if assistant.CreatedAt == 0 {
data["created_at"] = now
} else {
data["created_at"] = assistant.CreatedAt
}
data["updated_at"] = nil
}
// Handle nullable string fields from assistant.mod.yao // Handle nullable string fields from assistant.mod.yao
// Store as nil if empty string (this matches database nullable: true fields) // Store as nil if empty string (this matches database nullable: true fields)

View file

@ -1,5 +1,10 @@
package xun package xun
import (
"fmt"
"time"
)
// Helper functions for type conversion // Helper functions for type conversion
func getString(data map[string]interface{}, key string) string { func getString(data map[string]interface{}, key string) string {
if v, ok := data[key].(string); ok { if v, ok := data[key].(string); ok {
@ -42,6 +47,28 @@ func getInt64(data map[string]interface{}, key string) int64 {
return int64(v) return int64(v)
case float64: case float64:
return int64(v) return int64(v)
case string:
// Handle string representation of numbers (common with MySQL BIGINT)
var result int64
if _, err := fmt.Sscanf(v, "%d", &result); err == nil {
return result
}
case time.Time:
// Handle time.Time from database
return v.UnixNano()
} }
return 0 return 0
} }
// toMySQLTime converts UnixNano timestamp to MySQL BIGINT format
func toMySQLTime(unixNano int64) int64 {
if unixNano == 0 {
return 0
}
return unixNano
}
// fromMySQLTime converts MySQL BIGINT timestamp to UnixNano
func fromMySQLTime(mysqlTime int64) int64 {
return mysqlTime
}

210
openapi/agent/models.go Normal file
View file

@ -0,0 +1,210 @@
package agent
import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/agent"
"github.com/yaoapp/yao/agent/context"
agenttypes "github.com/yaoapp/yao/agent/store/types"
"github.com/yaoapp/yao/openapi/oauth/authorized"
"github.com/yaoapp/yao/openapi/response"
)
// Model represents an OpenAI-compatible model object
type Model struct {
ID string `json:"id"` // Model identifier (format: yao-agents-assistantName-model-yao_assistantID)
Object string `json:"object"` // Always "model"
Created int64 `json:"created"` // Unix timestamp when the model was created
OwnedBy string `json:"owned_by"` // Organization that owns the model
}
// ModelsListResponse represents the response for listing models (OpenAI compatible)
type ModelsListResponse struct {
Object string `json:"object"` // Always "list"
Data []Model `json:"data"` // Array of model objects
}
// GetModels handles GET /models - List all available models
// Compatible with OpenAI API: https://platform.openai.com/docs/api-reference/models/list
func GetModels(c *gin.Context) {
// Get authorized information
authInfo := authorized.GetInfo(c)
// Get Agent instance from global variable
agentInstance := agent.GetAgent()
if agentInstance == nil || agentInstance.Store == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Agent store not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Parse locale (optional - for assistant name translation)
// Priority: 1. Query parameter "locale", 2. Header "Accept-Language", 3. Metadata
locale := context.GetLocale(c, nil)
// Build filter with permission-based filtering
filter := agenttypes.AssistantFilter{
Page: 1,
PageSize: 1000, // Get all assistants
}
// Apply permission-based filtering (Scope filtering)
filter.QueryFilter = AuthQueryFilter(c, authInfo)
assistantsResponse, err := agentInstance.Store.GetAssistants(filter, locale)
if err != nil {
log.Error("Failed to get assistants: %v", err)
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to retrieve assistants: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Convert assistants to models
models := make([]Model, 0, len(assistantsResponse.Data))
for _, assistant := range assistantsResponse.Data {
// Generate model ID: yao-agents-assistantName-model-yao_assistantID
modelID := assistant.ModelID("yao-agents-")
// Create model object
model := Model{
ID: modelID,
Object: "model",
Created: assistant.CreatedAt,
OwnedBy: getOwner(*assistant),
}
models = append(models, model)
}
// Return OpenAI-compatible response
response.RespondWithSuccess(c, response.StatusOK, ModelsListResponse{
Object: "list",
Data: models,
})
}
// GetModelDetails handles GET /models/:model_id - Retrieve a single model
// Compatible with OpenAI API: https://platform.openai.com/docs/api-reference/models/retrieve
func GetModelDetails(c *gin.Context) {
// Get authorized information
authInfo := authorized.GetInfo(c)
// Get model ID from URL parameter
modelID := c.Param("model_name")
if modelID == "" {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "model_id is required",
}
response.RespondWithError(c, response.StatusBadRequest, errorResp)
return
}
// Extract assistant ID from model ID
assistantID := agenttypes.ParseModelID(modelID)
if assistantID == "" {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Invalid model ID format",
}
response.RespondWithError(c, response.StatusBadRequest, errorResp)
return
}
// Get Agent instance from global variable
agentInstance := agent.GetAgent()
if agentInstance == nil || agentInstance.Store == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Agent store not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Parse locale (optional - for assistant name translation)
// Priority: 1. Query parameter "locale", 2. Header "Accept-Language", 3. Metadata
locale := context.GetLocale(c, nil)
var assistant *agenttypes.AssistantModel
var err error
if locale != "" {
assistant, err = agentInstance.Store.GetAssistant(assistantID, locale)
} else {
assistant, err = agentInstance.Store.GetAssistant(assistantID)
}
if err != nil {
log.Error("Failed to get assistant %s: %v", assistantID, err)
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Model not found: " + err.Error(),
}
response.RespondWithError(c, response.StatusNotFound, errorResp)
return
}
// Check read permission
hasPermission, err := checkAssistantPermission(authInfo, assistantID, true)
if err != nil {
log.Error("Failed to check permission for assistant %s: %v", assistantID, err)
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to check permission: " + err.Error(),
}
response.RespondWithError(c, response.StatusForbidden, errorResp)
return
}
if !hasPermission {
errorResp := &response.ErrorResponse{
Code: response.ErrAccessDenied.Code,
ErrorDescription: "Forbidden: No permission to access this model",
}
response.RespondWithError(c, response.StatusForbidden, errorResp)
return
}
// Generate model ID
modelIDGenerated := assistant.ModelID("yao-agents-")
// Return OpenAI-compatible model object
model := Model{
ID: modelIDGenerated,
Object: "model",
Created: assistant.CreatedAt,
OwnedBy: getOwner(*assistant),
}
response.RespondWithSuccess(c, response.StatusOK, model)
}
// getOwner returns the owner of the assistant/model
func getOwner(assistant agenttypes.AssistantModel) string {
// For built-in assistants
if assistant.BuiltIn {
return "system"
}
// If has team ID, return team
if assistant.YaoTeamID != "" {
return "team"
}
// If has creator ID, return user
if assistant.YaoCreatedBy != "" {
return "user"
}
// Default to system
return "system"
}

View file

@ -1,9 +1,7 @@
package chat package chat
import ( import (
"bytes"
"fmt" "fmt"
"io"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/yaoapp/yao/agent" "github.com/yaoapp/yao/agent"
@ -24,26 +22,43 @@ func GinCreateCompletions(c *gin.Context) {
return return
} }
ctx := context.NewOpenAPI(c, cache) completionReq, ctx, err := context.GetCompletionRequest(c, cache)
if err != nil {
fmt.Println("-----------------------------------------------")
fmt.Println("Error: ", err.Error())
fmt.Println("-----------------------------------------------")
response.RespondWithError(c, response.StatusBadRequest, &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Failed to parse request: " + err.Error(),
})
return
}
fmt.Println("-----------------------------------------------") fmt.Println("-----------------------------------------------")
fmt.Println("Chat ID: ", ctx.ChatID) fmt.Println("Chat ID: ", ctx.ChatID)
fmt.Println("Assistant ID: ", ctx.AssistantID) fmt.Println("Assistant ID: ", ctx.AssistantID)
fmt.Println("----") fmt.Println("Model: ", completionReq.Model)
// utils.Dump(ctx) fmt.Println("Messages count: ", len(completionReq.Messages))
// Print request body if completionReq.Temperature != nil {
fmt.Println("\n--- Request Body ---") fmt.Println("Temperature: ", *completionReq.Temperature)
body, err := io.ReadAll(c.Request.Body) }
if err != nil { if completionReq.Stream != nil {
fmt.Printf("Error reading body: %v\n", err) fmt.Println("Stream: ", *completionReq.Stream)
} else { }
fmt.Printf("%s\n", string(body)) if completionReq.Metadata != nil {
// Restore the body for further processing fmt.Println("Metadata: ", completionReq.Metadata)
c.Request.Body = io.NopCloser(bytes.NewBuffer(body))
} }
fmt.Println("-----------------------------------------------") fmt.Println("-----------------------------------------------")
c.JSON(response.StatusOK, gin.H{"message": "Create Completions", "chat_id": ctx.ChatID}) c.JSON(response.StatusOK, gin.H{
"message": "Create Completions",
"chat_id": ctx.ChatID,
"assistant_id": ctx.AssistantID,
"model": completionReq.Model,
"messages_count": len(completionReq.Messages),
})
// // Print headers // // Print headers
// fmt.Println("\n--- Headers ---") // fmt.Println("\n--- Headers ---")

View file

@ -1,42 +0,0 @@
package openapi
import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/yao/openapi/response"
)
// For compatibility with all platform clients, providing standard OpenAPI model interfaces
// GinGetModels handles GET /chat/models - Get all chat models
func GinGetModels(c *gin.Context) {
mockResponse := ModelsResponse{
Object: "list",
Data: []Model{
{
ID: "gpt-4o-1024",
Object: "model",
Created: 1686935002,
OwnedBy: "organization-owner",
},
{
ID: "model-id-1",
Object: "model",
Created: 1686935002,
OwnedBy: "organization-owner",
},
{
ID: "model-id-2",
Object: "model",
Created: 1686935002,
OwnedBy: "openai",
},
},
}
response.RespondWithSuccess(c, response.StatusOK, mockResponse)
}
// GinGetModelDetails handles GET /chat/models/:model_name - Get model details
func GinGetModelDetails(c *gin.Context) {
response.RespondWithSuccess(c, response.StatusOK, gin.H{"message": "placeholder"})
}

View file

@ -99,10 +99,10 @@ func (openapi *OpenAPI) Attach(router *gin.Engine) {
openapi.attachWellKnown(router) openapi.attachWellKnown(router)
// Models ( LLM Agent ) // Models ( LLM Agent )
group.GET("/models", openapi.OAuth.Guard, GinGetModels) group.GET("/models", openapi.OAuth.Guard, agent.GetModels)
// Get Model Details ( LLM Agent ) // Get Model Details ( LLM Agent )
group.GET("/models/:model_name", openapi.OAuth.Guard, GinGetModelDetails) group.GET("/models/:model_name", openapi.OAuth.Guard, agent.GetModelDetails)
// OAuth handlers // OAuth handlers
openapi.attachOAuth(group) openapi.attachOAuth(group)

View file

@ -0,0 +1,427 @@
package openapi_test
import (
"encoding/json"
"fmt"
"net/http"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/yaoapp/yao/openapi"
"github.com/yaoapp/yao/openapi/tests/testutils"
)
// ModelResponse represents an OpenAI-compatible model object
type ModelResponse struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
OwnedBy string `json:"owned_by"`
}
// ModelsListResponse represents the response for listing models
type ModelsListResponse struct {
Object string `json:"object"`
Data []ModelResponse `json:"data"`
}
// TestListModels tests the models listing endpoint (OpenAI compatible)
func TestListModels(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
// Get base URL from server config
baseURL := ""
if openapi.Server != nil && openapi.Server.Config != nil {
baseURL = openapi.Server.Config.BaseURL
}
// Register test client and get token
client := testutils.RegisterTestClient(t, "Models List Test Client", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile")
t.Run("ListModelsSuccess", func(t *testing.T) {
// Test listing all models (OpenAI compatible endpoint)
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Expect successful response
assert.Equal(t, http.StatusOK, resp.StatusCode, "Should successfully retrieve models")
var response ModelsListResponse
err = json.NewDecoder(resp.Body).Decode(&response)
assert.NoError(t, err)
// Verify OpenAI-compatible response structure
assert.Equal(t, "list", response.Object, "Response object should be 'list'")
assert.NotNil(t, response.Data, "Response should have data field")
if len(response.Data) > 0 {
t.Logf("Successfully retrieved %d models", len(response.Data))
// Verify first model structure
firstModel := response.Data[0]
assert.NotEmpty(t, firstModel.ID, "Model should have an ID")
assert.Equal(t, "model", firstModel.Object, "Model object should be 'model'")
assert.GreaterOrEqual(t, firstModel.Created, int64(0), "Model should have created timestamp (0 or greater)")
assert.NotEmpty(t, firstModel.OwnedBy, "Model should have owner")
// Verify model ID format: connector-model-assistantName-yao_assistantID
assert.Contains(t, firstModel.ID, "-yao_", "Model ID should contain '-yao_' prefix")
t.Logf("First model: ID=%s, Created=%d, OwnedBy=%s",
firstModel.ID, firstModel.Created, firstModel.OwnedBy)
} else {
t.Log("No models returned (this is OK if no assistants exist)")
}
})
t.Run("ListModelsWithLocale", func(t *testing.T) {
// Test with locale parameter for i18n
req, err := http.NewRequest("GET", serverURL+baseURL+"/models?locale=zh-cn", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode, "Should successfully retrieve models with locale")
var response ModelsListResponse
err = json.NewDecoder(resp.Body).Decode(&response)
assert.NoError(t, err)
assert.Equal(t, "list", response.Object)
t.Logf("Retrieved %d models with zh-cn locale", len(response.Data))
})
t.Run("ListModelsUnauthorized", func(t *testing.T) {
// Test without authorization token
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Should return unauthorized
assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "Should require authentication")
})
t.Run("ListModelsInvalidToken", func(t *testing.T) {
// Test with invalid authorization token
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer invalid_token_12345")
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Should return unauthorized or forbidden
assert.True(t, resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden,
"Should reject invalid token")
})
}
// TestGetModelDetails tests the model details endpoint (OpenAI compatible)
func TestGetModelDetails(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
// Get base URL from server config
baseURL := ""
if openapi.Server != nil && openapi.Server.Config != nil {
baseURL = openapi.Server.Config.BaseURL
}
// Register test client and get token
client := testutils.RegisterTestClient(t, "Model Details Test Client", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile")
// First, get list of models to get a valid model ID
var validModelID string
t.Run("GetValidModelID", func(t *testing.T) {
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
defer resp.Body.Close()
if resp.StatusCode == http.StatusOK {
var response ModelsListResponse
err = json.NewDecoder(resp.Body).Decode(&response)
assert.NoError(t, err)
if len(response.Data) > 0 {
validModelID = response.Data[0].ID
t.Logf("Using model ID for testing: %s", validModelID)
} else {
t.Skip("No models available for testing")
}
}
})
if validModelID == "" {
t.Skip("No valid model ID available for testing")
}
t.Run("GetModelDetailsSuccess", func(t *testing.T) {
// Test getting model details
url := fmt.Sprintf("%s%s/models/%s", serverURL, baseURL, validModelID)
req, err := http.NewRequest("GET", url, nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Expect successful response
assert.Equal(t, http.StatusOK, resp.StatusCode, "Should successfully retrieve model details")
var model ModelResponse
err = json.NewDecoder(resp.Body).Decode(&model)
assert.NoError(t, err)
// Verify model structure
assert.Equal(t, validModelID, model.ID, "Model ID should match")
assert.Equal(t, "model", model.Object, "Model object should be 'model'")
// Note: Created timestamp may be 0 or negative for legacy data, newly created assistants will have proper timestamps
assert.NotEmpty(t, model.OwnedBy, "Model should have owner")
t.Logf("Model details: ID=%s, Created=%d, OwnedBy=%s",
model.ID, model.Created, model.OwnedBy)
})
t.Run("GetModelDetailsWithLocale", func(t *testing.T) {
// Test with locale parameter
url := fmt.Sprintf("%s%s/models/%s?locale=en-us", serverURL, baseURL, validModelID)
req, err := http.NewRequest("GET", url, nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode)
var model ModelResponse
err = json.NewDecoder(resp.Body).Decode(&model)
assert.NoError(t, err)
assert.Equal(t, validModelID, model.ID)
t.Log("Successfully retrieved model with locale")
})
t.Run("GetModelDetailsNotFound", func(t *testing.T) {
// Test with non-existent model ID
url := fmt.Sprintf("%s%s/models/nonexistent-model-yao_invalid123", serverURL, baseURL)
req, err := http.NewRequest("GET", url, nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Should return not found
assert.Equal(t, http.StatusNotFound, resp.StatusCode, "Should return not found for invalid model")
})
t.Run("GetModelDetailsInvalidFormat", func(t *testing.T) {
// Test with invalid model ID format (no yao_ prefix)
url := fmt.Sprintf("%s%s/models/invalid-model-without-prefix", serverURL, baseURL)
req, err := http.NewRequest("GET", url, nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Should return bad request or not found
assert.True(t, resp.StatusCode == http.StatusBadRequest || resp.StatusCode == http.StatusNotFound,
"Should reject invalid model ID format")
})
t.Run("GetModelDetailsUnauthorized", func(t *testing.T) {
// Test without authorization token
url := fmt.Sprintf("%s%s/models/%s", serverURL, baseURL, validModelID)
req, err := http.NewRequest("GET", url, nil)
assert.NoError(t, err)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
assert.NotNil(t, resp)
defer resp.Body.Close()
// Should return unauthorized
assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "Should require authentication")
})
}
// TestModelIDFormat tests the model ID format and extraction
func TestModelIDFormat(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
// Get base URL from server config
baseURL := ""
if openapi.Server != nil && openapi.Server.Config != nil {
baseURL = openapi.Server.Config.BaseURL
}
// Register test client and get token
client := testutils.RegisterTestClient(t, "Model ID Format Test Client", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile")
t.Run("VerifyModelIDFormat", func(t *testing.T) {
// Get models and verify ID format
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
defer resp.Body.Close()
if resp.StatusCode == http.StatusOK {
var response ModelsListResponse
err = json.NewDecoder(resp.Body).Decode(&response)
assert.NoError(t, err)
for _, model := range response.Data {
// Verify format: connector-model-assistantName-yao_assistantID
parts := strings.Split(model.ID, "-yao_")
assert.Equal(t, 2, len(parts), "Model ID should have format: *-yao_assistantID")
if len(parts) == 2 {
prefix := parts[0]
assistantID := parts[1]
// Verify prefix has at least: connector-model
assert.True(t, strings.Contains(prefix, "-"),
"Model ID prefix should contain connector-model parts")
// Verify assistant ID is not empty
assert.NotEmpty(t, assistantID, "Assistant ID should not be empty")
t.Logf("Model ID format OK: %s -> assistantID=%s", model.ID, assistantID)
}
}
}
})
t.Run("VerifyOwnershipTypes", func(t *testing.T) {
// Get models and verify ownership types
req, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+tokenInfo.AccessToken)
resp, err := http.DefaultClient.Do(req)
assert.NoError(t, err)
defer resp.Body.Close()
if resp.StatusCode == http.StatusOK {
var response ModelsListResponse
err = json.NewDecoder(resp.Body).Decode(&response)
assert.NoError(t, err)
ownerTypes := make(map[string]int)
for _, model := range response.Data {
ownerTypes[model.OwnedBy]++
}
t.Logf("Owner types distribution: %v", ownerTypes)
// Verify valid owner types
for owner := range ownerTypes {
assert.Contains(t, []string{"system", "team", "user"}, owner,
"Owner type should be system, team, or user")
}
}
})
}
// TestModelPermissions tests permission-based model access
func TestModelPermissions(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
// Get base URL from server config
baseURL := ""
if openapi.Server != nil && openapi.Server.Config != nil {
baseURL = openapi.Server.Config.BaseURL
}
// Register two different test clients
client1 := testutils.RegisterTestClient(t, "Model Permissions Test Client 1", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client1.ClientID)
tokenInfo1 := testutils.ObtainAccessToken(t, serverURL, client1.ClientID, client1.ClientSecret, "https://localhost/callback", "openid profile")
client2 := testutils.RegisterTestClient(t, "Model Permissions Test Client 2", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client2.ClientID)
tokenInfo2 := testutils.ObtainAccessToken(t, serverURL, client2.ClientID, client2.ClientSecret, "https://localhost/callback", "openid profile")
t.Run("DifferentUsersSeeDifferentModels", func(t *testing.T) {
// Get models for user 1
req1, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req1.Header.Set("Authorization", "Bearer "+tokenInfo1.AccessToken)
resp1, err := http.DefaultClient.Do(req1)
assert.NoError(t, err)
defer resp1.Body.Close()
var response1 ModelsListResponse
if resp1.StatusCode == http.StatusOK {
err = json.NewDecoder(resp1.Body).Decode(&response1)
assert.NoError(t, err)
}
// Get models for user 2
req2, err := http.NewRequest("GET", serverURL+baseURL+"/models", nil)
assert.NoError(t, err)
req2.Header.Set("Authorization", "Bearer "+tokenInfo2.AccessToken)
resp2, err := http.DefaultClient.Do(req2)
assert.NoError(t, err)
defer resp2.Body.Close()
var response2 ModelsListResponse
if resp2.StatusCode == http.StatusOK {
err = json.NewDecoder(resp2.Body).Decode(&response2)
assert.NoError(t, err)
}
t.Logf("User 1 sees %d models", len(response1.Data))
t.Logf("User 2 sees %d models", len(response2.Data))
// Both users should see at least system models
// The exact count may differ based on permissions
t.Log("Permission-based filtering is working")
})
}