package handler import ( "context" "encoding/base64" "encoding/json" "log" "net/http" "strings" "time" "github.com/alberdjuniawan/anput-api/middleware" "github.com/alberdjuniawan/anput-api/repository" "github.com/alberdjuniawan/anput-api/service" "github.com/alberdjuniawan/anput-api/util" chiMiddleware "github.com/go-chi/chi/v5/middleware" "github.com/gorilla/websocket" ) type WebSocketHandler struct { AIService service.AIService Repo repository.Repository Upgrader websocket.Upgrader } func NewWebSocketHandler(aiService service.AIService, repo repository.Repository) *WebSocketHandler { return &WebSocketHandler{ AIService: aiService, Repo: repo, Upgrader: websocket.Upgrader{ ReadBufferSize: 4096, WriteBufferSize: 4096, CheckOrigin: func(r *http.Request) bool { return true }, }, } } type wsResponse struct { Type string `json:"type"` Data interface{} `json:"data"` } func (h *WebSocketHandler) HandleStream(w http.ResponseWriter, r *http.Request) { reqID := chiMiddleware.GetReqID(r.Context()) userIDVal := r.Context().Value(middleware.UserIDKey) if userIDVal == nil { util.WriteError(w, http.StatusUnauthorized, "Unauthorized", reqID) return } userID := userIDVal.(string) slug := r.URL.Query().Get("slug") if slug == "" { util.WriteError(w, http.StatusBadRequest, "Missing slug", reqID) return } schemaData, err := h.Repo.GetSchemaBySlug(r.Context(), userID, slug) if err != nil { log.Printf("[%s] WS Schema Error: %v", reqID, err) util.WriteError(w, http.StatusNotFound, "Schema not found", reqID) return } conn, err := h.Upgrader.Upgrade(w, r, nil) if err != nil { log.Printf("[%s] WS Upgrade Failed: %v", reqID, err) return } defer conn.Close() log.Printf("[%s] WS Connected", reqID) var fullTranscriptBuilder strings.Builder sentenceBuf := make([]byte, 0) lastDataTime := time.Now() const bytesPerSec = 32000 minChunkSize := bytesPerSec / 2 silenceThreshold := 700 * time.Millisecond audioCh := make(chan []byte, 100) stopCh := make(chan bool) go func() { defer close(audioCh) for { mt, p, err := conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) { log.Printf("[%s] WS Read Error: %v", reqID, err) } stopCh <- true return } if mt == websocket.BinaryMessage { audioCh <- p } else if mt == websocket.TextMessage { var msg map[string]string if json.Unmarshal(p, &msg) == nil && msg["action"] == "stop" { stopCh <- true return } if b, err := base64.StdEncoding.DecodeString(string(p)); err == nil { audioCh <- b } } } }() ticker := time.NewTicker(200 * time.Millisecond) defer ticker.Stop() for { select { case chunk := <-audioCh: sentenceBuf = append(sentenceBuf, chunk...) lastDataTime = time.Now() case <-ticker.C: if len(sentenceBuf) > minChunkSize && time.Since(lastDataTime) > silenceThreshold { h.processChunk(conn, sentenceBuf, &fullTranscriptBuilder, reqID) sentenceBuf = make([]byte, 0) } case <-stopCh: if len(sentenceBuf) > 0 { h.processChunk(conn, sentenceBuf, &fullTranscriptBuilder, reqID) } finalText := fullTranscriptBuilder.String() log.Printf("[%s] Final Transcript: %s", reqID, finalText) if len(strings.TrimSpace(finalText)) > 2 { res, err := h.AIService.ExtractFromText(context.Background(), finalText, schemaData.Definition) if err == nil { conn.WriteJSON(wsResponse{ Type: "result", Data: map[string]interface{}{ "parsed": res.Parsed, "final_transcript": finalText, }, }) } else { log.Printf("[%s] Extraction Failed: %v", reqID, err) conn.WriteJSON(wsResponse{Type: "error", Data: "Extraction failed: " + err.Error()}) } } log.Printf("[%s] WS Session Ended", reqID) return } } } func (h *WebSocketHandler) processChunk(conn *websocket.Conn, audio []byte, sb *strings.Builder, reqID string) { text, err := h.AIService.TranscribeChunk(context.Background(), audio) if err == nil && len(strings.TrimSpace(text)) > 0 { clean := strings.TrimSpace(text) sb.WriteString(clean + " ") conn.WriteJSON(wsResponse{ Type: "transcript", Data: sb.String(), }) } else if err != nil { log.Printf("[%s] Transcribe Chunk Error: %v", reqID, err) } }