158 lines
5.2 KiB
Go
158 lines
5.2 KiB
Go
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
|
||
}
|