Files
agent/backend/internal/logic/ai/sessions.go

158 lines
5.2 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package ai
import (
"errors"
"fmt"
"strings"
"time"
"gorm.io/gorm"
"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")
ErrExpertNotFound = errors.New("ai expert not found")
)
type sessionGateway interface {
ReserveRateLimit(userID uint, action string, limit int, window time.Duration) error
SelectKey(userID uint) (SelectedKey, error)
RecordCall(database *gorm.DB, userID uint, provider string, usedKeyType string, action string, status string, errText string) error
}
// SessionService 只管理项目内普通 AI 会话;它不会把上下文自动转换为任务、笔记或资料。
type SessionService struct {
gateway sessionGateway
}
func NewSessionService(gateway sessionGateway) *SessionService {
return &SessionService{gateway: gateway}
}
// List 在项目 owner 校验后返回该项目的会话,内部自增 ID 不离开服务边界。
func (s *SessionService) List(userID uint, projectIdentity string) ([]models.SaAISession, error) {
project, err := projects.FindOwnedProject(userID, projectIdentity)
if err != nil {
return nil, err
}
var sessions []models.SaAISession
if err := models.DBService.Preload("Expert").Where("project_id = ?", project.ID).Order("updated_at desc, id desc").Find(&sessions).Error; err != nil {
return nil, err
}
if sessions == nil {
sessions = []models.SaAISession{}
}
return sessions, nil
}
// Create 先校验项目,再限流,最后才选择 provider/key会话创建不会生成任何正式业务对象。
func (s *SessionService) Create(userID uint, projectIdentity, title, context string) (*models.SaAISession, error) {
return s.CreateWithExpert(userID, projectIdentity, title, context, "")
}
// CreateWithExpert 创建带本地专家角色的项目会话。
func (s *SessionService) CreateWithExpert(userID uint, projectIdentity, title, context, expertIdentity string) (*models.SaAISession, 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")
}
expert, err := findExpertByIdentity(models.DBService, expertIdentity)
if err != nil {
return nil, err
}
if err := s.gateway.ReserveRateLimit(userID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour); err != nil {
if errors.Is(err, ErrAIRateLimited) {
if auditErr := s.gateway.RecordCall(models.DBService, 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(models.DBService, userID, "none", "none", aiSessionCreateAction, "failed", "rate_limit_reservation_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"
}
provider, keyType := selectedAuditMetadata(selected)
if auditErr := s.gateway.RecordCall(models.DBService, userID, provider, keyType, aiSessionCreateAction, "failed", code); auditErr != nil {
return nil, fmt.Errorf("select ai key: %v; record failure: %w", err, auditErr)
}
return nil, err
}
session := models.SaAISession{
ProjectID: project.ID,
CreatedBy: userID,
Title: title,
Context: context,
Status: defaultSessionStatus,
}
if expert != nil {
session.ExpertID = &expert.ID
session.ExpertIdentity = &expert.Identity
}
failureCode := "session_create_failed"
err = models.DBService.Transaction(func(tx *gorm.DB) error {
if err := tx.Create(&session).Error; err != nil {
return err
}
// ready 只表示会话入口已建立,不表示 provider 已回复或任何业务对象已创建。
failureCode = "audit_write_failed"
if err := s.gateway.RecordCall(tx, userID, selected.Provider, selected.KeyType, aiSessionCreateAction, defaultSessionStatus, ""); err != nil {
return err
}
failureCode = "session_transaction_failed"
return nil
})
if err != nil {
if auditErr := s.gateway.RecordCall(models.DBService, userID, selected.Provider, selected.KeyType, aiSessionCreateAction, "failed", failureCode); auditErr != nil {
return nil, fmt.Errorf("create ai session transaction: %v; record failure: %w", err, auditErr)
}
return nil, err
}
session.Expert = expert
return &session, nil
}
func selectedAuditMetadata(selected SelectedKey) (string, string) {
provider := strings.TrimSpace(selected.Provider)
keyType := strings.TrimSpace(selected.KeyType)
if provider == "" {
provider = "none"
}
if keyType == "" {
keyType = "none"
}
return provider, keyType
}
func aiSessionStatus(session models.SaAISession) string {
if status := strings.TrimSpace(session.Status); status != "" {
return status
}
return defaultSessionStatus
}