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.
385 lines
9.2 KiB
Go
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)
|
|
}
|