114 lines
3.7 KiB
Go
114 lines
3.7 KiB
Go
package ai
|
|
|
|
import (
|
|
"fmt"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
"senlinai-agent/backend/internal/models"
|
|
)
|
|
|
|
func TestSelectKeyPrefersUserKey(t *testing.T) {
|
|
database := newTestDB(t)
|
|
require.NoError(t, database.Create(&models.SenlinAgentAIKey{UserID: 3, Provider: "openai", EncryptedAPIKey: "user-key"}).Error)
|
|
gateway := NewGateway("system-key")
|
|
|
|
selected, err := gateway.SelectKey(3)
|
|
|
|
require.NoError(t, err)
|
|
require.Equal(t, "user", selected.KeyType)
|
|
require.Equal(t, "user-key", selected.APIKey)
|
|
}
|
|
|
|
func TestSaveUserKeyEncryptsStoredKey(t *testing.T) {
|
|
database := newTestDB(t)
|
|
gateway := NewGatewayWithSecret("system-key", "test-encryption-secret")
|
|
|
|
require.NoError(t, gateway.SaveUserKey(3, "openai", "user-key"))
|
|
|
|
var stored models.SenlinAgentAIKey
|
|
require.NoError(t, database.Where("user_id = ?", 3).First(&stored).Error)
|
|
require.NotEqual(t, "user-key", stored.EncryptedAPIKey)
|
|
require.Contains(t, stored.EncryptedAPIKey, "v1:")
|
|
selected, err := gateway.SelectKey(3)
|
|
require.NoError(t, err)
|
|
require.Equal(t, "user-key", selected.APIKey)
|
|
}
|
|
|
|
func TestSelectKeyFallsBackToSystemKey(t *testing.T) {
|
|
newTestDB(t)
|
|
gateway := NewGateway("system-key")
|
|
|
|
selected, err := gateway.SelectKey(3)
|
|
|
|
require.NoError(t, err)
|
|
require.Equal(t, "system", selected.KeyType)
|
|
require.Equal(t, "system-key", selected.APIKey)
|
|
}
|
|
|
|
func TestSelectKeyReturnsErrorWhenNoKeyAvailable(t *testing.T) {
|
|
newTestDB(t)
|
|
gateway := NewGateway("")
|
|
|
|
_, err := gateway.SelectKey(3)
|
|
|
|
require.ErrorContains(t, err, "no ai key available")
|
|
}
|
|
|
|
func TestRecordCallStoresAuditFields(t *testing.T) {
|
|
database := newTestDB(t)
|
|
gateway := NewGateway("system-key")
|
|
|
|
require.NoError(t, gateway.RecordCall(database, 3, "openai", "system", "inbox_analyze", "failed", "rate limited"))
|
|
|
|
var log models.SenlinAgentAICallLog
|
|
require.NoError(t, database.First(&log).Error)
|
|
require.Equal(t, uint(3), log.UserID)
|
|
require.Equal(t, "openai", log.Provider)
|
|
require.Equal(t, "system", log.UsedKeyType)
|
|
require.Equal(t, "inbox_analyze", log.Action)
|
|
require.Equal(t, "failed", log.Status)
|
|
require.Equal(t, "rate limited", log.Error)
|
|
}
|
|
|
|
func TestReserveRateLimitRejectsCallsOverWindow(t *testing.T) {
|
|
database := newTestDB(t)
|
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "rate@example.com", DisplayName: "Rate", PasswordHash: "hash"}).Error)
|
|
gateway := NewGateway("system-key")
|
|
require.NoError(t, gateway.ReserveRateLimit(1, "inbox_analyze", 1, time.Hour))
|
|
|
|
err := gateway.ReserveRateLimit(1, "inbox_analyze", 1, time.Hour)
|
|
|
|
require.ErrorContains(t, err, "ai rate limit exceeded")
|
|
}
|
|
|
|
func TestCreateAISession(t *testing.T) {
|
|
database := newTestDB(t)
|
|
user := models.SenlinAgentUser{Email: "session@example.com", DisplayName: "Session User", PasswordHash: "hash"}
|
|
require.NoError(t, database.Create(&user).Error)
|
|
project := models.SenlinAgentProject{OwnerID: user.ID, Name: "Session Project", Identifier: "SESSION"}
|
|
require.NoError(t, database.Create(&project).Error)
|
|
service := NewSessionService(NewGateway("system-key"))
|
|
|
|
session, err := service.Create(user.ID, project.Identity, "报价分析", "询价上下文")
|
|
|
|
require.NoError(t, err)
|
|
require.Equal(t, project.ID, session.ProjectID)
|
|
require.Equal(t, user.ID, session.CreatedBy)
|
|
require.Equal(t, "报价分析", session.Title)
|
|
require.Equal(t, "ready", session.Status)
|
|
}
|
|
|
|
func newTestDB(t *testing.T) *gorm.DB {
|
|
t.Helper()
|
|
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
|
require.NoError(t, err)
|
|
require.NoError(t, models.AutoMigrate(database))
|
|
models.DBService = database
|
|
return database
|
|
}
|