211 lines
6.6 KiB
Go
211 lines
6.6 KiB
Go
// Package service 提供 transcribe 模块的业务服务
|
||
package service
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"time"
|
||
|
||
meetingModel "github.com/echochat/backend/app/meeting/model"
|
||
"github.com/echochat/backend/app/transcribe/dao"
|
||
"github.com/echochat/backend/app/transcribe/model"
|
||
"github.com/echochat/backend/pkg/logs"
|
||
"go.uber.org/zap"
|
||
"gorm.io/gorm"
|
||
)
|
||
|
||
// 业务错误
|
||
var (
|
||
ErrRecordingNotFound = errors.New("录制不存在")
|
||
ErrRecordingNotReady = errors.New("录制尚未就绪,无法转写")
|
||
ErrTranscribeRunning = errors.New("已有转写任务进行中,请稍后再试")
|
||
ErrSTTNotConfigured = errors.New("语音转写服务未配置")
|
||
)
|
||
|
||
// TranscribeService 转写编排服务
|
||
//
|
||
// 设计要点:
|
||
// - Submit 立即返回 (transcript, 已 running),真正的 STT 调用走 goroutine
|
||
// - 同一 recording 只允许一个 running 任务,重复提交直接返回当前进行中行
|
||
// - 转写超时 10 分钟;超时后任务自标记 failed
|
||
// - 服务重启会留下 running 状态的孤儿行,留给后续兜底任务处理(暂不在本期实现)
|
||
type TranscribeService struct {
|
||
db *gorm.DB
|
||
transcripts *dao.TranscriptDAO
|
||
llmConfig *dao.LLMConfigDAO
|
||
sttClient *OpenAICompatibleClient
|
||
}
|
||
|
||
// NewTranscribeService 创建实例
|
||
func NewTranscribeService(
|
||
db *gorm.DB,
|
||
transcripts *dao.TranscriptDAO,
|
||
llmConfig *dao.LLMConfigDAO,
|
||
sttClient *OpenAICompatibleClient,
|
||
) *TranscribeService {
|
||
return &TranscribeService{
|
||
db: db,
|
||
transcripts: transcripts,
|
||
llmConfig: llmConfig,
|
||
sttClient: sttClient,
|
||
}
|
||
}
|
||
|
||
// IsAvailable 转写功能是否可用(LLM 配置库连通)
|
||
func (s *TranscribeService) IsAvailable() bool {
|
||
return s.llmConfig != nil && s.llmConfig.IsEnabled()
|
||
}
|
||
|
||
// Submit 提交转写任务
|
||
//
|
||
// 行为:
|
||
// - 录制必须存在且 status=ready
|
||
// - 如果已有 running 行,返回 ErrTranscribeRunning
|
||
// - 如果已有 ready 行且 force=false,直接返回已有结果
|
||
// - 否则:插入或重置 transcript 行 → 启动 goroutine 跑真实调用 → 同步返回 running 状态
|
||
//
|
||
// 参数:
|
||
// - force:true 表示强制重跑(覆盖 ready 行)
|
||
// - language:可选,"" 表示自动检测
|
||
//
|
||
// 返回值为 Submit 时刻的 transcript 快照,调用方可继续轮询 Get 拿最新状态。
|
||
func (s *TranscribeService) Submit(ctx context.Context, recordingID int64, force bool, language string) (*model.Transcript, error) {
|
||
if !s.IsAvailable() {
|
||
return nil, ErrSTTNotConfigured
|
||
}
|
||
|
||
// 1) 校验录制存在且已就绪
|
||
var rec meetingModel.MeetingRecording
|
||
err := s.db.WithContext(ctx).Where("id = ?", recordingID).First(&rec).Error
|
||
if err != nil {
|
||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||
return nil, ErrRecordingNotFound
|
||
}
|
||
return nil, err
|
||
}
|
||
if rec.Status != meetingModel.MeetingRecordingStatusReady || rec.FileURL == "" {
|
||
return nil, ErrRecordingNotReady
|
||
}
|
||
|
||
// 2) 检查现有 transcript
|
||
existing, err := s.transcripts.GetByRecordingID(ctx, recordingID)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if existing != nil {
|
||
switch existing.Status {
|
||
case model.TranscriptStatusRunning, model.TranscriptStatusPending:
|
||
return existing, ErrTranscribeRunning
|
||
case model.TranscriptStatusReady:
|
||
if !force {
|
||
return existing, nil
|
||
}
|
||
}
|
||
}
|
||
|
||
// 3) 创建/更新行为 pending(即将进入 running)
|
||
now := time.Now()
|
||
t := &model.Transcript{
|
||
RecordingID: recordingID,
|
||
RoomID: rec.RoomID,
|
||
Status: model.TranscriptStatusPending,
|
||
Segments: "[]",
|
||
StartedAt: &now,
|
||
}
|
||
if existing != nil {
|
||
t.ID = existing.ID
|
||
}
|
||
if err := s.transcripts.Upsert(ctx, t); err != nil {
|
||
return nil, fmt.Errorf("写入 transcript 失败: %w", err)
|
||
}
|
||
|
||
// 4) 取最新 ID(Upsert 后 t.ID 已填)
|
||
if err := s.transcripts.MarkRunning(ctx, t.ID); err != nil {
|
||
return nil, fmt.Errorf("切换 running 状态失败: %w", err)
|
||
}
|
||
|
||
// 5) 异步执行真实调用
|
||
// 使用全新 context,超时 10min;不沿用入参 ctx,避免 HTTP 请求结束后 ctx 被 cancel 中断后台任务
|
||
go s.runJob(t.ID, rec, language)
|
||
|
||
// 返回 running 状态快照
|
||
updated, _ := s.transcripts.GetByRecordingID(ctx, recordingID)
|
||
if updated != nil {
|
||
return updated, nil
|
||
}
|
||
return t, nil
|
||
}
|
||
|
||
// runJob goroutine 内执行真实 STT 调用
|
||
//
|
||
// 不返回错误:所有失败都写回数据库 status=failed + error_msg
|
||
func (s *TranscribeService) runJob(transcriptID int64, rec meetingModel.MeetingRecording, language string) {
|
||
const funcName = "TranscribeService.runJob"
|
||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute)
|
||
defer cancel()
|
||
|
||
logs.Info(ctx, funcName, "STT 任务开始",
|
||
zap.Int64("transcript_id", transcriptID),
|
||
zap.Int64("recording_id", rec.ID),
|
||
)
|
||
|
||
cfg, err := s.llmConfig.PickActiveSTT(ctx)
|
||
if err != nil {
|
||
s.fail(ctx, transcriptID, "选择 STT 配置失败: "+err.Error())
|
||
return
|
||
}
|
||
|
||
// 把 provider/model/key 元信息写回 transcript,便于审计
|
||
_ = s.db.WithContext(ctx).Model(&model.Transcript{}).
|
||
Where("id = ?", transcriptID).
|
||
Updates(map[string]any{
|
||
"provider_code": cfg.Provider.ProviderCode,
|
||
"model_code": cfg.Model.ModelCode,
|
||
"key_id": cfg.Key.ID,
|
||
}).Error
|
||
|
||
result, err := s.sttClient.Transcribe(ctx, cfg, rec.FileURL, language)
|
||
if err != nil {
|
||
s.fail(ctx, transcriptID, "STT 调用失败: "+err.Error())
|
||
return
|
||
}
|
||
|
||
segmentsJSON, err := json.Marshal(result.Segments)
|
||
if err != nil {
|
||
segmentsJSON = []byte("[]")
|
||
}
|
||
|
||
if err := s.transcripts.MarkReady(ctx, transcriptID,
|
||
result.Text, string(segmentsJSON),
|
||
result.Language, int(result.Duration),
|
||
); err != nil {
|
||
logs.Error(ctx, funcName, "写回 ready 状态失败", zap.Error(err))
|
||
return
|
||
}
|
||
logs.Info(ctx, funcName, "STT 任务完成",
|
||
zap.Int64("transcript_id", transcriptID),
|
||
zap.Int("text_len", len(result.Text)),
|
||
zap.Int("segments", len(result.Segments)),
|
||
)
|
||
}
|
||
|
||
// fail 统一失败收尾:写日志 + 落库
|
||
func (s *TranscribeService) fail(ctx context.Context, transcriptID int64, msg string) {
|
||
logs.Warn(ctx, "TranscribeService.fail", msg, zap.Int64("transcript_id", transcriptID))
|
||
if err := s.transcripts.MarkFailed(ctx, transcriptID, msg); err != nil {
|
||
logs.Error(ctx, "TranscribeService.fail", "写回 failed 状态失败", zap.Error(err))
|
||
}
|
||
}
|
||
|
||
// GetByRecording 拉取某段录制的转写记录(不存在返回 nil, nil)
|
||
func (s *TranscribeService) GetByRecording(ctx context.Context, recordingID int64) (*model.Transcript, error) {
|
||
return s.transcripts.GetByRecordingID(ctx, recordingID)
|
||
}
|
||
|
||
// ListByRoom 拉取一场会议下所有转写
|
||
func (s *TranscribeService) ListByRoom(ctx context.Context, roomID int64) ([]model.Transcript, error) {
|
||
return s.transcripts.ListByRoomID(ctx, roomID)
|
||
}
|