anput-api / handler /websocket.go
alberdjuniawan's picture
refactor: standardize API response
8a5f1ff
Raw History Blame Contribute Delete
4.42 kB
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)
}
}