diff --git a/internal/call/signaling.go b/internal/call/signaling.go index 4f36058..13da445 100644 --- a/internal/call/signaling.go +++ b/internal/call/signaling.go @@ -2,6 +2,7 @@ package call import ( "encoding/binary" + "io" "sync" "time" @@ -56,6 +57,8 @@ type Call struct { Active bool FromE2EEKey [32]byte ToE2EEKey [32]byte + TargetStream *quic.SendStream // server→target (forwards caller audio) + CallerStream *quic.SendStream // server→caller (forwards target audio) } func NewManager(logger *zap.Logger, sessions SessionRegistry, mediaTimeout time.Duration, brutalBPS uint64, applyBrutal BrutalApplier) *Manager { @@ -239,7 +242,7 @@ func (m *Manager) HandleAccept(sess CallSession, payload []byte) { } func (m *Manager) HandleReject(sess CallSession, _ []byte) { - m.endCall(sess, false) + m.endCall(sess, true) } func (m *Manager) HandleEnd(sess CallSession, _ []byte) { @@ -261,6 +264,12 @@ func (m *Manager) endCall(sess CallSession, notify bool) { delete(m.calls, callID) delete(m.userCalls, call.From) delete(m.userCalls, call.To) + if call.TargetStream != nil { + call.TargetStream.Close() + } + if call.CallerStream != nil { + call.CallerStream.Close() + } m.mu.Unlock() sess.ClearCall() @@ -291,28 +300,58 @@ func (m *Manager) GetCall(callID uint64) *Call { return m.calls[callID] } -func (m *Manager) RelayMedia(senderSess CallSession, data []byte) { +func writeMediaFrame(w io.Writer, data []byte) error { + header := make([]byte, 4) + binary.BigEndian.PutUint32(header, uint32(len(data))) + if _, err := w.Write(header); err != nil { + return err + } + _, err := w.Write(data) + return err +} + +func (m *Manager) RelayMediaToUniStream(senderSess CallSession, data []byte) { callID := senderSess.CallID() - m.mu.RLock() + m.mu.Lock() call, exists := m.calls[callID] - m.mu.RUnlock() if !exists || !call.Active { + m.mu.Unlock() return } - var target *quic.Conn + var target *quic.SendStream + var targetConn *quic.Conn if senderSess == call.FromSess { - target = call.ToSess.Conn() + target = call.TargetStream + targetConn = call.ToSess.Conn() } else if senderSess == call.ToSess { - target = call.FromSess.Conn() + target = call.CallerStream + targetConn = call.FromSess.Conn() } else { + m.mu.Unlock() return } - if err := target.SendDatagram(data); err != nil { - m.logger.Warn("Failed to relay media datagram", - zap.Uint64("call_id", callID), - zap.Error(err)) + if target == nil { + var err error + target, err = targetConn.OpenUniStream() + if err != nil { + m.logger.Warn("Failed to open uni stream for media relay", + zap.Uint64("call_id", callID), zap.Error(err)) + m.mu.Unlock() + return + } + if senderSess == call.FromSess { + call.TargetStream = target + } else { + call.CallerStream = target + } + } + m.mu.Unlock() + + if err := writeMediaFrame(target, data); err != nil { + m.logger.Warn("Failed to write media to uni stream", + zap.Uint64("call_id", callID), zap.Error(err)) } } diff --git a/internal/server/server.go b/internal/server/server.go index 8fca25f..8006e99 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -8,6 +8,7 @@ import ( "crypto/x509" "encoding/binary" "fmt" + "io" "sync" "time" @@ -167,6 +168,7 @@ func (s *Server) OnConnect(ctx context.Context, conn *quic.Conn) { go s.handleStreams(ctx, conn, sess) go s.handleDatagrams(ctx, conn, sess) + go s.handleUniStreams(ctx, conn, sess) <-ctx.Done() s.Unregister(sess.Number()) @@ -193,7 +195,42 @@ func (s *Server) handleDatagrams(ctx context.Context, conn *quic.Conn, sess *Ses continue } if s.callMgr != nil && sess.IsAuthenticated() && sess.CallID() != 0 { - s.callMgr.RelayMedia(sess, data) + s.callMgr.RelayMediaToUniStream(sess, data) + } + } +} + +func readMediaFrame(r io.Reader) ([]byte, error) { + header := make([]byte, 4) + if _, err := io.ReadFull(r, header); err != nil { + return nil, err + } + n := binary.BigEndian.Uint32(header) + data := make([]byte, n) + if _, err := io.ReadFull(r, data); err != nil { + return nil, err + } + return data, nil +} + +func (s *Server) handleUniStreams(ctx context.Context, conn *quic.Conn, sess *Session) { + for { + str, err := conn.AcceptUniStream(ctx) + if err != nil { + return + } + go s.handleAudioUniStream(ctx, str, sess) + } +} + +func (s *Server) handleAudioUniStream(ctx context.Context, stream *quic.ReceiveStream, sess *Session) { + for { + data, err := readMediaFrame(stream) + if err != nil { + return + } + if s.callMgr != nil && sess.IsAuthenticated() && sess.CallID() != 0 { + s.callMgr.RelayMediaToUniStream(sess, data) } } }