package stage import ( "context" "encoding/json" "fmt" "os" "path/filepath" "sort" "strings" "sync" "gitea.maximumdirect.net/eric/narratio/internal/adapters/whisperx" "gitea.maximumdirect.net/eric/narratio/internal/artifacts" "gitea.maximumdirect.net/eric/narratio/internal/manifest" ) type transcribeStage struct{} func (transcribeStage) Name() string { return "transcribe" } func (transcribeStage) Declares() IODecl { return IODecl{ Inputs: []artifacts.Ref{ {Kind: "audio", Category: "audio", RelativePath: "audio/*.flac"}, }, Outputs: []artifacts.Ref{ {Kind: "transcript_raw", Category: "transcripts", RelativePath: "transcripts/raw/*.json"}, }, } } func (transcribeStage) Run(ctx context.Context, env *Env, m *manifest.Manifest) (*StageResult, error) { if env == nil || env.Config == nil { return nil, fmt.Errorf("transcribe: stage environment config is required") } if env.ArtifactStore == nil { return nil, fmt.Errorf("transcribe: artifact store is required") } if env.Config.Pipeline == nil || env.Config.Session == nil { return nil, fmt.Errorf("transcribe: resolved config must include pipeline and session") } if env.WhisperX == nil { return nil, fmt.Errorf("transcribe: whisperx adapter is required") } var sessionID string if m != nil { sessionID = strings.TrimSpace(m.SessionID) } if sessionID == "" { sessionID = strings.TrimSpace(env.Config.Session.SessionID) } if sessionID == "" { return nil, fmt.Errorf("transcribe: session id is required") } paths := sessionPathsForEnv(env, sessionID) audioFiles, err := discoverPreparedAudio(m, paths.AudioDir) if err != nil { return nil, fmt.Errorf("transcribe: resolve audio inputs: %w", err) } if len(audioFiles) == 0 { return nil, fmt.Errorf("transcribe: no prepared audio files found") } concurrency := 1 if env.Config.Pipeline.WhisperX.Concurrency != nil && *env.Config.Pipeline.WhisperX.Concurrency > 0 { concurrency = *env.Config.Pipeline.WhisperX.Concurrency } if concurrency > len(audioFiles) { concurrency = len(audioFiles) } type job struct { speakerID string audioPath string outPath string } jobs := make([]job, 0, len(audioFiles)) seenSpeaker := map[string]string{} for _, audioPath := range audioFiles { base := strings.TrimSuffix(filepath.Base(audioPath), filepath.Ext(audioPath)) if base == "" { return nil, fmt.Errorf("transcribe: could not derive speaker id from %q", audioPath) } if prev, ok := seenSpeaker[base]; ok { return nil, fmt.Errorf("transcribe: duplicate speaker/audio basename %q from %q and %q", base, prev, audioPath) } seenSpeaker[base] = audioPath jobs = append(jobs, job{ speakerID: base, audioPath: audioPath, outPath: filepath.Join(paths.TranscriptsRawDir, base+".json"), }) } stageCtx, cancel := context.WithCancel(ctx) defer cancel() var ( wg sync.WaitGroup jobCh = make(chan job) mu sync.Mutex firstErr error perFile = map[string]map[string]any{} outputRef = map[string]artifacts.Ref{} ) worker := func() { defer wg.Done() for j := range jobCh { if stageCtx.Err() != nil { return } res, runErr := env.WhisperX.Transcribe(stageCtx, whisperx.TranscribeRequest{ SpeakerID: j.speakerID, AudioPath: j.audioPath, OutputRawTranscriptPath: j.outPath, }) if runErr != nil { mu.Lock() if firstErr == nil { firstErr = fmt.Errorf("speaker %q (%s): %w", j.speakerID, filepath.Base(j.audioPath), runErr) cancel() } mu.Unlock() continue } if strings.TrimSpace(res.OutputRawTranscriptPath) == "" { res.OutputRawTranscriptPath = j.outPath } if filepath.Clean(res.OutputRawTranscriptPath) != filepath.Clean(j.outPath) { mu.Lock() if firstErr == nil { firstErr = fmt.Errorf("speaker %q: adapter output path %q did not match expected %q", j.speakerID, res.OutputRawTranscriptPath, j.outPath) cancel() } mu.Unlock() continue } if err := validateTranscriptJSONFile(j.outPath); err != nil { mu.Lock() if firstErr == nil { firstErr = fmt.Errorf("speaker %q output %q invalid: %w", j.speakerID, j.outPath, err) cancel() } mu.Unlock() continue } meta := map[string]any{ "attempts": res.Attempts, "http_status": res.HTTPStatus, "duration_ms": res.Duration.Milliseconds(), "output_path": j.outPath, } mu.Lock() perFile[j.speakerID] = meta outputRef[j.speakerID] = artifacts.Ref{ Kind: "transcript_raw", Category: "transcripts", SessionID: sessionID, AbsolutePath: j.outPath, } mu.Unlock() } } wg.Add(concurrency) for i := 0; i < concurrency; i++ { go worker() } dispatch: for _, j := range jobs { select { case <-stageCtx.Done(): break dispatch case jobCh <- j: } } close(jobCh) wg.Wait() if firstErr != nil { return nil, firstErr } speakers := make([]string, 0, len(perFile)) for speaker := range perFile { speakers = append(speakers, speaker) } sort.Strings(speakers) outputs := make([]artifacts.Ref, 0, len(speakers)) outputPaths := make([]string, 0, len(speakers)) orderedPerFile := make(map[string]any, len(speakers)) for _, speaker := range speakers { ref := outputRef[speaker] outputs = append(outputs, ref) outputPaths = append(outputPaths, ref.AbsolutePath) orderedPerFile[speaker] = perFile[speaker] } retries := 0 if env.Config.Pipeline.WhisperX.Retries != nil { retries = *env.Config.Pipeline.WhisperX.Retries } return &StageResult{ Outputs: outputs, Metadata: map[string]any{ "stage": "transcribe", "audio_files_count": len(audioFiles), "language": env.Config.Pipeline.WhisperX.Language, "concurrency": concurrency, "retries": retries, "retry_delay": env.Config.Pipeline.WhisperX.RetryDelay, "timeout": env.Config.Pipeline.WhisperX.Timeout, "output_paths": outputPaths, "per_file": orderedPerFile, }, }, nil } func discoverPreparedAudio(m *manifest.Manifest, audioDir string) ([]string, error) { fromManifest := make([]string, 0) if m != nil { for _, in := range m.Inputs { if in.Kind != "audio" { continue } p := strings.TrimSpace(in.Path) if p == "" { continue } fromManifest = append(fromManifest, filepath.Clean(p)) } } if len(fromManifest) > 0 { sort.Strings(fromManifest) if err := validateAudioFiles(fromManifest); err != nil { return nil, err } return fromManifest, nil } entries, err := os.ReadDir(audioDir) if err != nil { return nil, fmt.Errorf("read audio directory %q: %w", audioDir, err) } out := make([]string, 0) for _, entry := range entries { if entry.IsDir() { continue } p := filepath.Join(audioDir, entry.Name()) if !isFlac(p) { continue } out = append(out, p) } sort.Strings(out) if err := validateAudioFiles(out); err != nil { return nil, err } return out, nil } func validateAudioFiles(paths []string) error { if len(paths) == 0 { return nil } seen := map[string]struct{}{} for _, p := range paths { clean := filepath.Clean(strings.TrimSpace(p)) if clean == "" { return fmt.Errorf("audio path is required") } if !isFlac(clean) { return fmt.Errorf("audio file %q must have .flac extension", clean) } info, err := os.Stat(clean) if err != nil { return fmt.Errorf("audio file %q not found: %w", clean, err) } if !info.Mode().IsRegular() { return fmt.Errorf("audio file %q is not a regular file", clean) } if _, exists := seen[clean]; exists { return fmt.Errorf("duplicate audio file path %q", clean) } seen[clean] = struct{}{} } return nil } func validateTranscriptJSONFile(path string) error { if err := requireFile(path, "raw transcript"); err != nil { return err } data, err := os.ReadFile(path) if err != nil { return fmt.Errorf("read transcript: %w", err) } var payload any if err := json.Unmarshal(data, &payload); err != nil { return fmt.Errorf("decode json: %w", err) } return nil }