qcc/internal/server/server.go
Niko Marmeladkov 0ed9d18b27 Fix nil interface trap in GetByNumber
Map lookup returns typed nil *Session for missing keys, which wraps into CallSession interface as non-nil typed nil. Nil check silently passes, causing nil dereference.
2026-06-30 14:33:41 +03:00

385 lines
9.2 KiB
Go

package server
import (
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/tls"
"encoding/binary"
"fmt"
"sync"
"github.com/apernet/quic-go"
"github.com/caddyserver/certmagic"
"github.com/libdns/cloudflare"
"go.uber.org/zap"
"github.com/niko/qcc/internal/auth"
"github.com/niko/qcc/internal/ca"
"github.com/niko/qcc/internal/call"
"github.com/niko/qcc/internal/congestion/brutal"
"github.com/niko/qcc/internal/config"
"github.com/niko/qcc/internal/identity"
"github.com/niko/qcc/internal/protocol"
"github.com/niko/qcc/internal/store"
"github.com/niko/qcc/internal/transport"
"github.com/niko/qcc/pkg/types"
)
type Server struct {
cfg *config.Config
logger *zap.Logger
st *store.Engine
ca *ca.CA
identMgr *identity.Manager
callMgr *call.Manager
trans *transport.Server
sessions map[string]*Session
sessionsMu sync.RWMutex
cancel context.CancelFunc
}
func New(cfg *config.Config, logger *zap.Logger) (*Server, error) {
st, err := store.NewEngine(cfg.DataDir)
if err != nil {
return nil, fmt.Errorf("init store: %w", err)
}
masterKey := make([]byte, 32)
if _, err := rand.Read(masterKey); err != nil {
return nil, fmt.Errorf("generate master key: %w", err)
}
var mk [32]byte
copy(mk[:], masterKey)
ca, err := ca.NewOrLoad(cfg.DataDir, cfg.CA.KeyType, cfg.CA.Validity)
if err != nil {
return nil, fmt.Errorf("init ca: %w", err)
}
identMgr, err := identity.NewManager(st, mk, cfg.Identity.Prefix, cfg.Identity.Cooldown, logger)
if err != nil {
return nil, fmt.Errorf("init identity manager: %w", err)
}
return &Server{
cfg: cfg,
logger: logger,
st: st,
ca: ca,
identMgr: identMgr,
sessions: make(map[string]*Session),
}, nil
}
func (s *Server) Start(ctx context.Context) error {
ctx, s.cancel = context.WithCancel(ctx)
tlsCfg, err := s.setupTLS()
if err != nil {
return fmt.Errorf("tls setup: %w", err)
}
brutalBPS, err := parseBandwidth(s.cfg.Relay.BandwidthUp)
if err != nil {
s.logger.Warn("Failed to parse relay bandwidth, using default", zap.Error(err))
brutalBPS = 10 * 1000 * 1000 / 8 // 10 mbps default
}
s.callMgr = call.NewManager(
s.logger,
s,
s.cfg.Relay.MediaTimeout,
brutalBPS,
func(conn *quic.Conn, bps uint64) {
conn.SetCongestionControl(brutal.NewBrutalSender(bps))
},
)
s.trans, err = transport.NewServer(s.cfg.Server.Listen, tlsCfg, s.logger, s)
if err != nil {
return fmt.Errorf("transport: %w", err)
}
err = s.trans.Start(ctx)
s.shutdown()
return err
}
func (s *Server) shutdown() {
s.logger.Info("Shutting down server...")
if s.callMgr != nil {
s.callMgr.Stop()
}
if s.trans != nil {
s.trans.Close()
}
s.sessionsMu.Lock()
for number, sess := range s.sessions {
if conn := sess.Conn(); conn != nil {
conn.CloseWithError(0, "server shutdown")
}
delete(s.sessions, number)
}
s.sessionsMu.Unlock()
if s.st != nil {
s.st.Close()
}
s.logger.Info("Server shutdown complete")
}
func (s *Server) setupTLS() (*tls.Config, error) {
if len(s.cfg.ACME.Domains) == 0 || s.cfg.ACME.Email == "" {
s.logger.Warn("ACME not configured, starting without TLS (insecure)")
return &tls.Config{
MinVersion: tls.VersionTLS13,
NextProtos: []string{"qcc"},
InsecureSkipVerify: true,
}, nil
}
dnsSolver := &certmagic.DNS01Solver{
DNSProvider: &cloudflare.Provider{
APIToken: s.cfg.ACME.DNS.Config["cloudflare_api_token"],
},
}
return transport.NewTLSConfig(
s.cfg.ACME.Domains,
s.cfg.ACME.Email,
s.cfg.ACME.CA,
dnsSolver,
s.logger,
)
}
func (s *Server) OnConnect(ctx context.Context, conn *quic.Conn) {
sess := NewSession(conn)
s.logger.Info("New connection",
zap.String("remote", conn.RemoteAddr().String()))
go s.handleStreams(ctx, conn, sess)
go s.handleDatagrams(ctx, conn, sess)
<-ctx.Done()
s.Unregister(sess.Number())
conn.CloseWithError(0, "shutdown")
}
func (s *Server) handleStreams(ctx context.Context, conn *quic.Conn, sess *Session) {
for {
str, err := conn.AcceptStream(ctx)
if err != nil {
return
}
go s.handleStream(ctx, str, sess)
}
}
func (s *Server) handleDatagrams(ctx context.Context, conn *quic.Conn, sess *Session) {
for {
data, err := conn.ReceiveDatagram(ctx)
if err != nil {
return
}
if len(data) == 0 {
continue
}
if s.callMgr != nil && sess.IsAuthenticated() && sess.CallID() != 0 {
s.callMgr.RelayMedia(conn, data)
}
}
}
func (s *Server) handleStream(ctx context.Context, stream *quic.Stream, sess *Session) {
st := transport.NewStream(stream)
sess.SetStream(st)
for {
frame, err := st.ReadFrame()
if err != nil {
return
}
switch types.OpCode(frame.OpCode) {
case types.OpGetChallenge:
s.handleGetChallenge(sess)
case types.OpSolve:
s.handleSolve(sess, frame.Payload)
case types.OpReroll:
s.handleReroll(sess)
case types.OpDial:
if !sess.IsAuthenticated() {
sendOpError(st, types.ErrInvalidRequest, "not authenticated")
continue
}
s.callMgr.HandleDial(sess, frame.Payload)
case types.OpAccept:
s.callMgr.HandleAccept(sess, frame.Payload)
case types.OpReject:
s.callMgr.HandleReject(sess, frame.Payload)
case types.OpEnd:
s.callMgr.HandleEnd(sess, frame.Payload)
default:
sendOpError(st, types.ErrInvalidRequest, "unknown opcode")
}
}
}
func (s *Server) handleGetChallenge(sess *Session) {
challenge := auth.NewChallenge(s.cfg.PoW.Difficulty)
sess.SetChallenge(challenge)
if st := sess.Stream(); st != nil {
st.WriteFrame(byte(types.OpChallenge), challenge.Marshal())
}
}
func (s *Server) handleSolve(sess *Session, payload []byte) {
solution := auth.UnmarshalSolution(payload)
if solution == nil {
sendOpError(sess.Stream(), types.ErrPoWInvalid, "invalid solution format")
return
}
challenge := sess.GetChallenge()
if challenge == nil {
sendOpError(sess.Stream(), types.ErrPoWInvalid, "no challenge issued")
return
}
if err := challenge.Verify(solution, s.cfg.PoW.ChallengeTTL); err != nil {
sendOpError(sess.Stream(), types.ErrPoWInvalid, err.Error())
return
}
var pubKey [32]byte
copy(pubKey[:], solution.ClientPubKey[:])
ident, err := s.identMgr.Allocate(pubKey)
if err != nil {
sendOpError(sess.Stream(), types.ErrInternal, err.Error())
return
}
certDER, err := s.ca.IssueCert(ident.Number, ed25519.PublicKey(pubKey[:]), s.cfg.CA.Validity)
if err != nil {
sendOpError(sess.Stream(), types.ErrInternal, "cert issuance failed")
return
}
ident.CertDER = certDER
sess.Authenticate(ident.Number, pubKey, certDER)
s.Register(sess, ident.Number)
payload = marshalIdentityResponse(ident.Number, certDER, s.ca.CACertPEM())
if st := sess.Stream(); st != nil {
st.WriteFrame(byte(types.OpIdentity), payload)
}
s.logger.Info("Client authenticated",
zap.String("number", ident.Number))
}
func (s *Server) handleReroll(sess *Session) {
if !sess.IsAuthenticated() {
sendOpError(sess.Stream(), types.ErrInvalidRequest, "not authenticated")
return
}
newIdent, err := s.identMgr.Reroll(sess.PubKey())
if err != nil {
if _, ok := err.(*identity.CooldownError); ok {
sendOpError(sess.Stream(), types.ErrCooldown, err.Error())
} else {
sendOpError(sess.Stream(), types.ErrInternal, err.Error())
}
return
}
certDER, err := s.ca.IssueCert(newIdent.Number, ed25519.PublicKey(newIdent.PubKey[:]), s.cfg.CA.Validity)
if err != nil {
sendOpError(sess.Stream(), types.ErrInternal, "cert issuance failed")
return
}
newIdent.CertDER = certDER
s.Unregister(sess.Number())
sess.Authenticate(newIdent.Number, newIdent.PubKey, certDER)
s.Register(sess, newIdent.Number)
payload := marshalIdentityResponse(newIdent.Number, certDER, s.ca.CACertPEM())
if st := sess.Stream(); st != nil {
st.WriteFrame(byte(types.OpIdentity), payload)
}
s.logger.Info("Number rerolled",
zap.String("old", sess.Number()),
zap.String("new", newIdent.Number))
}
func marshalIdentityResponse(number string, certDER, caCertPEM []byte) []byte {
payload := make([]byte, len(number)+1+2+len(certDER)+len(caCertPEM))
off := 0
copy(payload[off:], []byte(number))
off += len(number)
payload[off] = 0
off++
binary.BigEndian.PutUint16(payload[off:off+2], uint16(len(certDER)))
off += 2
copy(payload[off:], certDER)
off += len(certDER)
copy(payload[off:], caCertPEM)
return payload
}
func (s *Server) Register(sess call.CallSession, number string) {
s.sessionsMu.Lock()
defer s.sessionsMu.Unlock()
s.sessions[number] = sess.(*Session)
}
func (s *Server) Unregister(number string) {
if number == "" {
return
}
s.sessionsMu.Lock()
defer s.sessionsMu.Unlock()
delete(s.sessions, number)
}
func (s *Server) GetByNumber(number string) call.CallSession {
s.sessionsMu.RLock()
defer s.sessionsMu.RUnlock()
sess := s.sessions[number]
if sess == nil {
return nil
}
return sess
}
func (s *Server) GetAll() []call.CallSession {
s.sessionsMu.RLock()
defer s.sessionsMu.RUnlock()
result := make([]call.CallSession, 0, len(s.sessions))
for _, sess := range s.sessions {
result = append(result, sess)
}
return result
}
func sendOpError(st protocol.FrameReadWriter, code types.ErrorCode, msg string) {
if st == nil {
return
}
payload := make([]byte, 2+len(msg))
binary.BigEndian.PutUint16(payload[0:2], uint16(code))
copy(payload[2:], msg)
st.WriteFrame(byte(types.OpError), payload)
}