feat: align controlled AI sessions and MVP controls
This commit is contained in:
@@ -10,6 +10,7 @@ import (
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
@@ -25,6 +26,11 @@ type SelectedKey struct {
|
||||
KeyType string
|
||||
}
|
||||
|
||||
var (
|
||||
ErrAIKeyMissing = errors.New("no ai key available")
|
||||
ErrAIRateLimited = errors.New("ai rate limit exceeded")
|
||||
)
|
||||
|
||||
func NewGateway(systemKey string) *Gateway {
|
||||
return NewGatewayWithSecret(systemKey, "development-ai-key-secret-change-me")
|
||||
}
|
||||
@@ -47,15 +53,19 @@ func (g *Gateway) SaveUserKey(userID uint, provider string, apiKey string) error
|
||||
|
||||
func (g *Gateway) SelectKey(userID uint) (SelectedKey, error) {
|
||||
var userKey models.SenlinAgentAIKey
|
||||
if err := models.DBService.Where("user_id = ?", userID).First(&userKey).Error; err == nil {
|
||||
err := models.DBService.Where("user_id = ?", userID).First(&userKey).Error
|
||||
if err == nil {
|
||||
apiKey, err := decryptAPIKey(userKey.EncryptedAPIKey, g.encryptionSecret)
|
||||
if err != nil {
|
||||
return SelectedKey{}, err
|
||||
}
|
||||
return SelectedKey{Provider: userKey.Provider, APIKey: apiKey, KeyType: "user"}, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return SelectedKey{}, err
|
||||
}
|
||||
if g.systemKey == "" {
|
||||
return SelectedKey{}, errors.New("no ai key available")
|
||||
return SelectedKey{}, ErrAIKeyMissing
|
||||
}
|
||||
return SelectedKey{Provider: "openai", APIKey: g.systemKey, KeyType: "system"}, nil
|
||||
}
|
||||
@@ -82,7 +92,7 @@ func (g *Gateway) CheckRateLimit(userID uint, action string, limit int, window t
|
||||
return err
|
||||
}
|
||||
if count >= int64(limit) {
|
||||
return errors.New("ai rate limit exceeded")
|
||||
return ErrAIRateLimited
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
@@ -85,20 +86,25 @@ func TestCheckRateLimitRejectsCallsOverWindow(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCreateAISession(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewSessionService()
|
||||
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(7, 3, "报价分析")
|
||||
session, err := service.Create(user.ID, project.Identity, "报价分析", "询价上下文")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uint(7), session.ProjectID)
|
||||
require.Equal(t, uint(3), session.CreatedBy)
|
||||
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{})
|
||||
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
|
||||
|
||||
120
backend/internal/logic/ai/handlers.go
Normal file
120
backend/internal/logic/ai/handlers.go
Normal file
@@ -0,0 +1,120 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/httpx"
|
||||
"senlinai-agent/backend/internal/logic/auth"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
// SessionDTO 是普通项目 AI 会话的公开契约,不携带数据库主键或自动创建对象的 ID。
|
||||
type SessionDTO struct {
|
||||
ID string `json:"id"`
|
||||
ProjectID string `json:"projectId"`
|
||||
Title string `json:"title"`
|
||||
Context string `json:"context"`
|
||||
Status string `json:"status"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
}
|
||||
|
||||
type createSessionRequest struct {
|
||||
Title string `json:"title"`
|
||||
Context string `json:"context"`
|
||||
}
|
||||
|
||||
// Handler 注册受认证、受项目所有权保护的 AI 会话接口。
|
||||
type Handler struct {
|
||||
service *SessionService
|
||||
}
|
||||
|
||||
func NewHandler(service *SessionService) *Handler {
|
||||
return &Handler{service: service}
|
||||
}
|
||||
|
||||
func (h *Handler) Register(router gin.IRouter) {
|
||||
router.GET("/projects/:projectId/ai-sessions", h.list)
|
||||
router.POST("/projects/:projectId/ai-sessions", h.create)
|
||||
}
|
||||
|
||||
func (h *Handler) list(c *gin.Context) {
|
||||
userID, projectIdentity, ok := aiRequestContext(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
sessions, err := h.service.List(userID, projectIdentity)
|
||||
if err != nil {
|
||||
writeAIError(c, err)
|
||||
return
|
||||
}
|
||||
items := make([]SessionDTO, 0, len(sessions))
|
||||
for _, session := range sessions {
|
||||
items = append(items, sessionDTO(session))
|
||||
}
|
||||
c.JSON(http.StatusOK, items)
|
||||
}
|
||||
|
||||
func (h *Handler) create(c *gin.Context) {
|
||||
userID, projectIdentity, ok := aiRequestContext(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input createSessionRequest
|
||||
if err := c.ShouldBindJSON(&input); err != nil {
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请求参数无效")
|
||||
return
|
||||
}
|
||||
session, err := h.service.Create(userID, projectIdentity, input.Title, input.Context)
|
||||
if err != nil {
|
||||
writeAIError(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, sessionDTO(*session))
|
||||
}
|
||||
|
||||
func aiRequestContext(c *gin.Context) (uint, string, bool) {
|
||||
userID, ok := auth.CurrentUserID(c)
|
||||
if !ok {
|
||||
httpx.Error(c, http.StatusUnauthorized, "unauthorized", "未登录或登录已失效")
|
||||
return 0, "", false
|
||||
}
|
||||
projectIdentity, ok := httpx.IdentityParam(c, "projectId")
|
||||
if !ok {
|
||||
return 0, "", false
|
||||
}
|
||||
return userID, projectIdentity, true
|
||||
}
|
||||
|
||||
func sessionDTO(session models.SenlinAgentAISession) SessionDTO {
|
||||
return SessionDTO{
|
||||
ID: session.Identity,
|
||||
ProjectID: session.ProjectIdentity,
|
||||
Title: session.Title,
|
||||
Context: session.Context,
|
||||
Status: aiSessionStatus(session),
|
||||
CreatedAt: session.CreatedAt.UTC(),
|
||||
UpdatedAt: session.UpdatedAt.UTC(),
|
||||
}
|
||||
}
|
||||
|
||||
func writeAIError(c *gin.Context, err error) {
|
||||
switch {
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
httpx.Error(c, http.StatusNotFound, "not_found", "项目不存在或无权访问")
|
||||
case errors.Is(err, ErrInvalidSession):
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请输入 AI 会话标题")
|
||||
case errors.Is(err, ErrAIRateLimited):
|
||||
httpx.Error(c, http.StatusTooManyRequests, "ai_rate_limited", "AI 请求过于频繁,请稍后重试")
|
||||
case errors.Is(err, ErrAIKeyMissing):
|
||||
httpx.Error(c, http.StatusServiceUnavailable, "ai_key_missing", "尚未配置可用的 AI 密钥")
|
||||
default:
|
||||
log.Printf("ai session request failed: %v", err)
|
||||
httpx.Error(c, http.StatusInternalServerError, "internal_error", "AI 会话操作失败,请稍后重试")
|
||||
}
|
||||
}
|
||||
278
backend/internal/logic/ai/handlers_test.go
Normal file
278
backend/internal/logic/ai/handlers_test.go
Normal file
@@ -0,0 +1,278 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
"senlinai-agent/backend/internal/config"
|
||||
"senlinai-agent/backend/internal/httpx"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
func TestAISessionHandlersRequireOwnedProject(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
intruder := createAIHandlerUser(t, database, "intruder@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "PRIVATE")
|
||||
gateway := &recordingSessionGateway{}
|
||||
router := aiHandlerTestRouter(intruder.ID, gateway)
|
||||
|
||||
for _, request := range []*http.Request{
|
||||
authenticatedAIRequest(t, http.MethodGet, "/api/v1/projects/"+project.Identity+"/ai-sessions", nil),
|
||||
authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "越权会话", "context": "不得创建",
|
||||
}),
|
||||
} {
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusNotFound, recorder.Code)
|
||||
var payload httpx.ErrorEnvelope
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
require.Equal(t, "not_found", payload.Error.Code)
|
||||
require.Equal(t, "项目不存在或无权访问", payload.Error.Message)
|
||||
}
|
||||
|
||||
var count int64
|
||||
require.NoError(t, database.Model(&models.SenlinAgentAISession{}).Count(&count).Error)
|
||||
require.Zero(t, count)
|
||||
require.Empty(t, gateway.steps)
|
||||
}
|
||||
|
||||
func TestCreateAISessionChecksRateLimitBeforeSelectingProvider(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "RATE")
|
||||
gateway := &recordingSessionGateway{
|
||||
selected: SelectedKey{Provider: "openai", APIKey: "system-key", KeyType: "system"},
|
||||
}
|
||||
router := aiHandlerTestRouter(owner.ID, gateway)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "限流顺序", "context": "只创建受控会话",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusCreated, recorder.Code)
|
||||
require.Equal(t, []string{"rate", "select", "record"}, gateway.steps)
|
||||
}
|
||||
|
||||
func TestCreateAISessionReturnsRateLimitBeforeMissingKeyAndAuditsFailure(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "LIMITED")
|
||||
gateway := NewGatewayWithSecret("", "test-encryption-secret")
|
||||
for range aiSessionCreateLimit {
|
||||
require.NoError(t, gateway.RecordCall(owner.ID, "openai", "system", "ai_session_create", "ready", ""))
|
||||
}
|
||||
router := aiHandlerTestRouter(owner.ID, gateway)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "超过限额", "context": "必须先返回限流",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusTooManyRequests, recorder.Code)
|
||||
var payload httpx.ErrorEnvelope
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
require.Equal(t, "ai_rate_limited", payload.Error.Code)
|
||||
require.Equal(t, "AI 请求过于频繁,请稍后重试", payload.Error.Message)
|
||||
var latest models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.Order("id desc").First(&latest).Error)
|
||||
require.Equal(t, "none", latest.Provider)
|
||||
require.Equal(t, "none", latest.UsedKeyType)
|
||||
require.Equal(t, "ai_session_create", latest.Action)
|
||||
require.Equal(t, "failed", latest.Status)
|
||||
require.Equal(t, "ai_rate_limited", latest.Error)
|
||||
}
|
||||
|
||||
func TestCreateAISessionWithoutKeyReturnsAuditedErrorAndCreatesNoFormalObjects(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "NO_KEY")
|
||||
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("", "test-encryption-secret"))
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "缺少密钥", "context": "不得伪装为已完成",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
|
||||
var payload httpx.ErrorEnvelope
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
require.Equal(t, "ai_key_missing", payload.Error.Code)
|
||||
require.Equal(t, "尚未配置可用的 AI 密钥", payload.Error.Message)
|
||||
for _, model := range []any{
|
||||
&models.SenlinAgentAISession{},
|
||||
&models.SenlinAgentTask{},
|
||||
&models.SenlinAgentNote{},
|
||||
&models.SenlinAgentSource{},
|
||||
} {
|
||||
var count int64
|
||||
require.NoError(t, database.Model(model).Count(&count).Error)
|
||||
require.Zero(t, count)
|
||||
}
|
||||
var call models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.First(&call).Error)
|
||||
require.Equal(t, owner.ID, call.UserID)
|
||||
require.Equal(t, "none", call.Provider)
|
||||
require.Equal(t, "none", call.UsedKeyType)
|
||||
require.Equal(t, "ai_session_create", call.Action)
|
||||
require.Equal(t, "failed", call.Status)
|
||||
require.Equal(t, "ai_key_missing", call.Error)
|
||||
}
|
||||
|
||||
func TestCreateAISessionReturnsIdentityDTOAndCompleteAuditWithoutAutomaticObjectIDs(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "CREATE")
|
||||
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("system-key", "test-encryption-secret"))
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "报价分析", "context": "仅整理会话上下文",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusCreated, recorder.Code)
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
require.ElementsMatch(t, []string{"id", "projectId", "title", "context", "status", "createdAt", "updatedAt"}, aiMapKeys(payload))
|
||||
identity, err := uuid.Parse(payload["id"].(string))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uuid.Version(7), identity.Version())
|
||||
require.Equal(t, project.Identity, payload["projectId"])
|
||||
require.Equal(t, "报价分析", payload["title"])
|
||||
require.Equal(t, "仅整理会话上下文", payload["context"])
|
||||
require.Equal(t, "ready", payload["status"])
|
||||
for _, forbidden := range []string{"taskId", "noteId", "sourceId", "createdTaskId", "createdNoteId", "createdSourceId"} {
|
||||
require.NotContains(t, payload, forbidden)
|
||||
}
|
||||
|
||||
var call models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.First(&call).Error)
|
||||
require.Equal(t, "openai", call.Provider)
|
||||
require.Equal(t, "system", call.UsedKeyType)
|
||||
require.Equal(t, "ai_session_create", call.Action)
|
||||
require.Equal(t, "ready", call.Status)
|
||||
require.Empty(t, call.Error)
|
||||
for _, model := range []any{&models.SenlinAgentTask{}, &models.SenlinAgentNote{}, &models.SenlinAgentSource{}} {
|
||||
var count int64
|
||||
require.NoError(t, database.Model(model).Count(&count).Error)
|
||||
require.Zero(t, count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListAISessionsReturnsOnlyOwnedProjectIdentityDTOs(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "owner@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "LIST")
|
||||
otherProject := createAIHandlerProject(t, database, owner.ID, "OTHER")
|
||||
require.NoError(t, database.Create(&models.SenlinAgentAISession{
|
||||
ProjectID: project.ID, CreatedBy: owner.ID, Title: "目标会话", Context: "项目上下文",
|
||||
}).Error)
|
||||
require.NoError(t, database.Create(&models.SenlinAgentAISession{
|
||||
ProjectID: otherProject.ID, CreatedBy: owner.ID, Title: "其他会话", Context: "不得混入",
|
||||
}).Error)
|
||||
router := aiHandlerTestRouter(owner.ID, &recordingSessionGateway{})
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodGet, "/api/v1/projects/"+project.Identity+"/ai-sessions", nil))
|
||||
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
var payload []map[string]any
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
require.Len(t, payload, 1)
|
||||
require.ElementsMatch(t, []string{"id", "projectId", "title", "context", "status", "createdAt", "updatedAt"}, aiMapKeys(payload[0]))
|
||||
require.Equal(t, project.Identity, payload[0]["projectId"])
|
||||
require.Equal(t, "目标会话", payload[0]["title"])
|
||||
require.Equal(t, "ready", payload[0]["status"])
|
||||
}
|
||||
|
||||
type recordingSessionGateway struct {
|
||||
steps []string
|
||||
selected SelectedKey
|
||||
rateErr error
|
||||
selectErr error
|
||||
}
|
||||
|
||||
func (g *recordingSessionGateway) CheckRateLimit(uint, string, int, time.Duration) error {
|
||||
g.steps = append(g.steps, "rate")
|
||||
return g.rateErr
|
||||
}
|
||||
|
||||
func (g *recordingSessionGateway) SelectKey(uint) (SelectedKey, error) {
|
||||
g.steps = append(g.steps, "select")
|
||||
return g.selected, g.selectErr
|
||||
}
|
||||
|
||||
func (g *recordingSessionGateway) RecordCall(uint, string, string, string, string, string) error {
|
||||
g.steps = append(g.steps, "record")
|
||||
return nil
|
||||
}
|
||||
|
||||
func newAIHandlerTestDB(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
|
||||
}
|
||||
|
||||
func createAIHandlerUser(t *testing.T, database *gorm.DB, email string) models.SenlinAgentUser {
|
||||
t.Helper()
|
||||
user := models.SenlinAgentUser{Email: email, DisplayName: email, PasswordHash: "hash"}
|
||||
require.NoError(t, database.Create(&user).Error)
|
||||
return user
|
||||
}
|
||||
|
||||
func createAIHandlerProject(t *testing.T, database *gorm.DB, ownerID uint, identifier string) models.SenlinAgentProject {
|
||||
t.Helper()
|
||||
project := models.SenlinAgentProject{OwnerID: ownerID, Name: identifier, Identifier: identifier}
|
||||
require.NoError(t, database.Create(&project).Error)
|
||||
return project
|
||||
}
|
||||
|
||||
func aiHandlerTestRouter(userID uint, gateway sessionGateway) http.Handler {
|
||||
return httpx.NewProtectedRouter(
|
||||
config.Config{Env: "test"},
|
||||
func(string) (uint, error) { return userID, nil },
|
||||
NewHandler(NewSessionService(gateway)),
|
||||
)
|
||||
}
|
||||
|
||||
func authenticatedAIRequest(t *testing.T, method, path string, body any) *http.Request {
|
||||
t.Helper()
|
||||
var requestBody *bytes.Reader
|
||||
if body == nil {
|
||||
requestBody = bytes.NewReader(nil)
|
||||
} else {
|
||||
encoded, err := json.Marshal(body)
|
||||
require.NoError(t, err)
|
||||
requestBody = bytes.NewReader(encoded)
|
||||
}
|
||||
request := httptest.NewRequest(method, path, requestBody)
|
||||
request.Header.Set("Authorization", "Bearer test-token")
|
||||
if body != nil {
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
return request
|
||||
}
|
||||
|
||||
func aiMapKeys(values map[string]any) []string {
|
||||
keys := make([]string, 0, len(values))
|
||||
for key := range values {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
@@ -2,23 +2,118 @@ package ai
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"senlinai-agent/backend/internal/logic/projects"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
const (
|
||||
aiSessionCreateAction = "ai_session_create"
|
||||
aiSessionCreateLimit = 20
|
||||
defaultSessionStatus = "ready"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidSession = errors.New("invalid ai session")
|
||||
)
|
||||
|
||||
type sessionGateway interface {
|
||||
CheckRateLimit(userID uint, action string, limit int, window time.Duration) error
|
||||
SelectKey(userID uint) (SelectedKey, error)
|
||||
RecordCall(userID uint, provider string, usedKeyType string, action string, status string, errText string) error
|
||||
}
|
||||
|
||||
// SessionService 只管理项目内普通 AI 会话;它不会把上下文自动转换为任务、笔记或资料。
|
||||
type SessionService struct {
|
||||
gateway sessionGateway
|
||||
}
|
||||
|
||||
func NewSessionService() *SessionService {
|
||||
return &SessionService{}
|
||||
func NewSessionService(gateway sessionGateway) *SessionService {
|
||||
return &SessionService{gateway: gateway}
|
||||
}
|
||||
|
||||
func (s *SessionService) Create(projectID uint, userID uint, title string) (*models.SenlinAgentAISession, error) {
|
||||
title = strings.TrimSpace(title)
|
||||
if title == "" {
|
||||
return nil, errors.New("session title is required")
|
||||
// List 在项目 owner 校验后返回该项目的会话,内部自增 ID 不离开服务边界。
|
||||
func (s *SessionService) List(userID uint, projectIdentity string) ([]models.SenlinAgentAISession, error) {
|
||||
project, err := projects.FindOwnedProject(userID, projectIdentity)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
session := &models.SenlinAgentAISession{ProjectID: projectID, CreatedBy: userID, Title: title}
|
||||
return session, models.DBService.Create(session).Error
|
||||
var sessions []models.SenlinAgentAISession
|
||||
if err := models.DBService.Where("project_id = ?", project.ID).Order("updated_at desc, id desc").Find(&sessions).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if sessions == nil {
|
||||
sessions = []models.SenlinAgentAISession{}
|
||||
}
|
||||
return sessions, nil
|
||||
}
|
||||
|
||||
// Create 先校验项目,再限流,最后才选择 provider/key;会话创建不会生成任何正式业务对象。
|
||||
func (s *SessionService) Create(userID uint, projectIdentity, title, context string) (*models.SenlinAgentAISession, error) {
|
||||
title = strings.TrimSpace(title)
|
||||
context = strings.TrimSpace(context)
|
||||
if title == "" {
|
||||
return nil, ErrInvalidSession
|
||||
}
|
||||
project, err := projects.FindOwnedProject(userID, projectIdentity)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if s.gateway == nil {
|
||||
return nil, errors.New("ai gateway is required")
|
||||
}
|
||||
|
||||
if err := s.gateway.CheckRateLimit(userID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour); err != nil {
|
||||
if errors.Is(err, ErrAIRateLimited) {
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", "ai_rate_limited"); auditErr != nil {
|
||||
return nil, fmt.Errorf("record ai rate limit failure: %w", auditErr)
|
||||
}
|
||||
return nil, ErrAIRateLimited
|
||||
}
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", "rate_limit_check_failed"); auditErr != nil {
|
||||
return nil, fmt.Errorf("check rate limit: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
selected, err := s.gateway.SelectKey(userID)
|
||||
if err != nil {
|
||||
code := "provider_selection_failed"
|
||||
if errors.Is(err, ErrAIKeyMissing) {
|
||||
code = "ai_key_missing"
|
||||
}
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", code); auditErr != nil {
|
||||
return nil, fmt.Errorf("select ai key: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
session := models.SenlinAgentAISession{
|
||||
ProjectID: project.ID,
|
||||
CreatedBy: userID,
|
||||
Title: title,
|
||||
Context: context,
|
||||
Status: defaultSessionStatus,
|
||||
}
|
||||
if err := models.DBService.Create(&session).Error; err != nil {
|
||||
if auditErr := s.gateway.RecordCall(userID, selected.Provider, selected.KeyType, aiSessionCreateAction, "failed", "session_create_failed"); auditErr != nil {
|
||||
return nil, fmt.Errorf("create ai session: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
// ready 只表示会话入口已建立,不表示 provider 已回复或任何业务对象已创建。
|
||||
if err := s.gateway.RecordCall(userID, selected.Provider, selected.KeyType, aiSessionCreateAction, defaultSessionStatus, ""); err != nil {
|
||||
return nil, fmt.Errorf("record ai session creation: %w", err)
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
func aiSessionStatus(session models.SenlinAgentAISession) string {
|
||||
if status := strings.TrimSpace(session.Status); status != "" {
|
||||
return status
|
||||
}
|
||||
return defaultSessionStatus
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user