mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-30 11:31:29 +02:00
Add peer-to-peer file drop
Files move directly between peers over the overlay, with no server in the path. The receiver listens on the WireGuard address only, so the port is unreachable from outside the tunnel, and every offer is matched to a known peer before anything is read. Consent is the default: an offer carries metadata alone, and no payload moves until the receiver accepts. Policy is per profile and device-local — off, ask, or auto-accept, with per-sender exceptions on top. Policy and history live in the profile's preferences, so removing a profile takes its file drop state with it. Transfers interrupted by a restart are settled on load; nothing survives to finish them, and left alone they would sit in the log as permanently pending. The Android bindings pull payload bytes through a chunk-returning stream: gomobile copies a []byte argument into a fresh Java array and never copies it back, so a fill-my-buffer method would hand back the right length with no data.
This commit is contained in:
@@ -27,6 +27,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/iface/device"
|
||||
"github.com/netbirdio/netbird/client/iface/netstack"
|
||||
"github.com/netbirdio/netbird/client/internal/dns"
|
||||
"github.com/netbirdio/netbird/client/internal/filedrop"
|
||||
"github.com/netbirdio/netbird/client/internal/lazyconn"
|
||||
"github.com/netbirdio/netbird/client/internal/listener"
|
||||
"github.com/netbirdio/netbird/client/internal/metrics"
|
||||
@@ -64,10 +65,11 @@ type ConnectClient struct {
|
||||
config *profilemanager.Config
|
||||
statusRecorder *peer.Status
|
||||
|
||||
engine *Engine
|
||||
engineMutex sync.Mutex
|
||||
clientMetrics *metrics.ClientMetrics
|
||||
updateManager *updater.Manager
|
||||
engine *Engine
|
||||
engineMutex sync.Mutex
|
||||
clientMetrics *metrics.ClientMetrics
|
||||
updateManager *updater.Manager
|
||||
fileDropManager *filedrop.Manager
|
||||
|
||||
persistSyncResponse bool
|
||||
}
|
||||
@@ -95,6 +97,12 @@ func (c *ConnectClient) SetUpdateManager(um *updater.Manager) {
|
||||
c.updateManager = um
|
||||
}
|
||||
|
||||
// SetFileDropManager hands the engine the active profile's file drop manager, so
|
||||
// the transfer receiver starts and stops with the tunnel. Must be set before Run.
|
||||
func (c *ConnectClient) SetFileDropManager(m *filedrop.Manager) {
|
||||
c.fileDropManager = m
|
||||
}
|
||||
|
||||
// Run with main logic.
|
||||
func (c *ConnectClient) Run(runningChan chan struct{}, logPath string) error {
|
||||
if androidRunOverride != nil {
|
||||
@@ -424,6 +432,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
||||
UpdateManager: c.updateManager,
|
||||
ClientMetrics: c.clientMetrics,
|
||||
MetricsCtx: c.ctx,
|
||||
FileDrop: c.fileDropManager,
|
||||
}, mobileDependency)
|
||||
engine.SetSyncResponsePersistence(c.persistSyncResponse)
|
||||
c.engine = engine
|
||||
|
||||
@@ -40,6 +40,7 @@ import (
|
||||
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
|
||||
"github.com/netbirdio/netbird/client/internal/dnsfwd"
|
||||
"github.com/netbirdio/netbird/client/internal/expose"
|
||||
"github.com/netbirdio/netbird/client/internal/filedrop"
|
||||
"github.com/netbirdio/netbird/client/internal/ingressgw"
|
||||
"github.com/netbirdio/netbird/client/internal/lazyconn"
|
||||
"github.com/netbirdio/netbird/client/internal/metrics"
|
||||
@@ -181,6 +182,7 @@ type EngineServices struct {
|
||||
UpdateManager *updater.Manager
|
||||
ClientMetrics *metrics.ClientMetrics
|
||||
MetricsCtx context.Context
|
||||
FileDrop *filedrop.Manager
|
||||
}
|
||||
|
||||
// Engine is a mechanism responsible for reacting on Signal and Management stream events and managing connections to the remote peers.
|
||||
@@ -236,6 +238,10 @@ type Engine struct {
|
||||
|
||||
sshServer sshServer
|
||||
|
||||
fileDrop *filedrop.Manager
|
||||
fileDropRunning bool
|
||||
fileDropPort uint16
|
||||
|
||||
statusRecorder *peer.Status
|
||||
|
||||
firewall firewallManager.Manager
|
||||
@@ -350,6 +356,7 @@ func NewEngine(
|
||||
metricsCtx: services.MetricsCtx,
|
||||
updateManager: services.UpdateManager,
|
||||
syncStoreDir: config.StateDir,
|
||||
fileDrop: services.FileDrop,
|
||||
}
|
||||
// sessionWatcher keeps the SubscribeStatus consumers in sync with the
|
||||
// session expiry deadline. Deadline-change ticks come for free via
|
||||
@@ -415,6 +422,8 @@ func (e *Engine) stopLocked() {
|
||||
log.Warnf("failed to stop SSH server: %v", err)
|
||||
}
|
||||
|
||||
e.stopFileDrop()
|
||||
|
||||
e.cleanupSSHConfig()
|
||||
|
||||
if e.ingressGatewayMgr != nil {
|
||||
@@ -1293,6 +1302,8 @@ func (e *Engine) updateConfig(conf *mgmProto.PeerConfig) error {
|
||||
}
|
||||
}
|
||||
|
||||
e.startFileDrop()
|
||||
|
||||
state := e.statusRecorder.GetLocalPeerState()
|
||||
state.IP = e.wgInterface.Address().String()
|
||||
state.IPv6 = e.wgInterface.Address().IPv6String()
|
||||
@@ -1960,6 +1971,8 @@ func (e *Engine) receiveSignalEvents() error {
|
||||
return err
|
||||
}
|
||||
|
||||
e.recordFiledropPort(msg.Key, msg.GetBody().GetFiledropPort())
|
||||
|
||||
log.Debugf("receiveMSG: took %s to get lock for peer %s with session id %s", gotLock, msg.Key, offerAnswer.SessionID)
|
||||
|
||||
if msg.Body.Type == sProto.Body_OFFER {
|
||||
|
||||
127
client/internal/engine_filedrop.go
Normal file
127
client/internal/engine_filedrop.go
Normal file
@@ -0,0 +1,127 @@
|
||||
package internal
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/netip"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/filedrop"
|
||||
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
)
|
||||
|
||||
type filedropResolver struct {
|
||||
status *peer.Status
|
||||
}
|
||||
|
||||
// ResolvePeer implements filedrop.PeerResolver.
|
||||
func (r filedropResolver) ResolvePeer(addr netip.Addr) (filedrop.PeerKey, string, bool) {
|
||||
state, ok := r.status.PeerStateByIP(addr.String())
|
||||
if !ok {
|
||||
return "", "", false
|
||||
}
|
||||
return filedrop.PeerKey(state.PubKey), state.FQDN, true
|
||||
}
|
||||
|
||||
func (e *Engine) startFileDrop() {
|
||||
if e.fileDrop == nil || e.fileDropRunning || e.wgInterface == nil {
|
||||
return
|
||||
}
|
||||
if e.config.BlockInbound {
|
||||
log.Info("file drop receiver is disabled because inbound connections are blocked")
|
||||
e.setFileDropTunnel()
|
||||
return
|
||||
}
|
||||
|
||||
wgAddr := e.wgInterface.Address()
|
||||
addr := netip.AddrPortFrom(wgAddr.IP, filedrop.Port)
|
||||
resolver := filedropResolver{status: e.statusRecorder}
|
||||
|
||||
netstackNet := e.wgInterface.GetNet()
|
||||
if err := e.fileDrop.StartReceiver(e.ctx, addr, netstackNet, resolver); err != nil {
|
||||
log.Errorf("failed to start file drop receiver: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
bound := e.fileDrop.ReceiverPort()
|
||||
if bound == 0 {
|
||||
bound = filedrop.Port
|
||||
}
|
||||
e.fileDropPort = bound
|
||||
|
||||
if v6 := wgAddr.IPv6; v6.IsValid() {
|
||||
if err := e.fileDrop.AddReceiverListener(e.ctx, netip.AddrPortFrom(v6, bound)); err != nil {
|
||||
log.Warnf("failed to add IPv6 file drop listener: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if netstackNet != nil {
|
||||
if registrar, ok := e.firewall.(interface {
|
||||
RegisterNetstackService(protocol nftypes.Protocol, port uint16)
|
||||
}); ok {
|
||||
registrar.RegisterNetstackService(nftypes.TCP, bound)
|
||||
}
|
||||
}
|
||||
|
||||
if bound != filedrop.Port {
|
||||
e.signaler.SetFiledropPort(bound)
|
||||
}
|
||||
|
||||
e.setFileDropTunnel()
|
||||
e.fileDropRunning = true
|
||||
}
|
||||
|
||||
// recordFiledropPort stores the file drop port a peer advertised over signaling;
|
||||
// a value that does not fit a port is treated as the default.
|
||||
func (e *Engine) recordFiledropPort(peerKey string, port uint32) {
|
||||
if e.fileDrop == nil {
|
||||
return
|
||||
}
|
||||
if port > 65535 {
|
||||
port = 0
|
||||
}
|
||||
e.fileDrop.Ports().Set(filedrop.PeerKey(peerKey), uint16(port))
|
||||
}
|
||||
|
||||
func (e *Engine) setFileDropTunnel() {
|
||||
var dial filedrop.DialFunc
|
||||
if netstackNet := e.wgInterface.GetNet(); netstackNet != nil {
|
||||
dial = func(ctx context.Context, _, addr string) (net.Conn, error) {
|
||||
addrPort, err := netip.ParseAddrPort(addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return netstackNet.DialContextTCPAddrPort(ctx, addrPort)
|
||||
}
|
||||
} else {
|
||||
dialer := &net.Dialer{}
|
||||
dial = dialer.DialContext
|
||||
}
|
||||
|
||||
e.fileDrop.SetTunnel(dial, e.statusRecorder.GetLocalPeerState().FQDN)
|
||||
}
|
||||
|
||||
func (e *Engine) stopFileDrop() {
|
||||
if e.fileDrop == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if e.fileDropRunning {
|
||||
if netstackNet := e.wgInterface.GetNet(); netstackNet != nil {
|
||||
if registrar, ok := e.firewall.(interface {
|
||||
UnregisterNetstackService(protocol nftypes.Protocol, port uint16)
|
||||
}); ok {
|
||||
registrar.UnregisterNetstackService(nftypes.TCP, e.fileDropPort)
|
||||
}
|
||||
}
|
||||
e.signaler.SetFiledropPort(0)
|
||||
}
|
||||
|
||||
if err := e.fileDrop.StopReceiver(); err != nil {
|
||||
log.Warnf("failed to stop file drop receiver: %v", err)
|
||||
}
|
||||
e.fileDropRunning = false
|
||||
e.fileDropPort = 0
|
||||
}
|
||||
446
client/internal/filedrop/client.go
Normal file
446
client/internal/filedrop/client.go
Normal file
@@ -0,0 +1,446 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultPollTimeout = 60 * time.Second
|
||||
defaultOfferTimeout = DefaultOfferTTL
|
||||
uploadRetryDelay = 2 * time.Second
|
||||
maxUploadAttempts = 3
|
||||
)
|
||||
|
||||
// DialFunc opens a connection to the receiving peer over the tunnel.
|
||||
type DialFunc func(ctx context.Context, network, addr string) (net.Conn, error)
|
||||
|
||||
// Payload is one item to send; Open is called per attempt starting at an offset.
|
||||
type Payload struct {
|
||||
Meta FileMeta
|
||||
Open func(offset int64) (io.ReadCloser, error)
|
||||
}
|
||||
|
||||
// ProgressFunc reports staged bytes for one item as the upload streams.
|
||||
type ProgressFunc func(index int, sent int64, total int64)
|
||||
|
||||
// ClientConfig configures the sending side.
|
||||
type ClientConfig struct {
|
||||
Dial DialFunc
|
||||
SenderName string
|
||||
PollTimeout time.Duration
|
||||
OfferTimeout time.Duration
|
||||
}
|
||||
|
||||
type progressReader struct {
|
||||
r io.Reader
|
||||
sent int64
|
||||
total int64
|
||||
report func(sent int64)
|
||||
}
|
||||
|
||||
// Client sends offers and payloads to a peer's file drop service.
|
||||
type Client struct {
|
||||
http *http.Client
|
||||
senderName string
|
||||
pollTimeout time.Duration
|
||||
offerTimeout time.Duration
|
||||
}
|
||||
|
||||
func (p *progressReader) Read(b []byte) (int, error) {
|
||||
n, err := p.r.Read(b)
|
||||
if n > 0 {
|
||||
p.sent += int64(n)
|
||||
p.report(p.sent)
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
// NewClient builds a sending client over the given dialer.
|
||||
func NewClient(cfg ClientConfig) (*Client, error) {
|
||||
if cfg.Dial == nil {
|
||||
return nil, errors.New("dial function is required")
|
||||
}
|
||||
|
||||
pollTimeout := cfg.PollTimeout
|
||||
if pollTimeout <= 0 {
|
||||
pollTimeout = defaultPollTimeout
|
||||
}
|
||||
offerTimeout := cfg.OfferTimeout
|
||||
if offerTimeout <= 0 {
|
||||
offerTimeout = defaultOfferTimeout
|
||||
}
|
||||
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { return cfg.Dial(ctx, network, addr) },
|
||||
MaxIdleConnsPerHost: 2,
|
||||
ResponseHeaderTimeout: pollTimeout + 30*time.Second,
|
||||
}
|
||||
|
||||
return &Client{
|
||||
http: &http.Client{Transport: transport},
|
||||
senderName: cfg.SenderName,
|
||||
pollTimeout: pollTimeout,
|
||||
offerTimeout: offerTimeout,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// TextPayload builds an inline text payload, which is carried in the offer itself.
|
||||
func TextPayload(name, text string) Payload {
|
||||
return Payload{
|
||||
Meta: FileMeta{
|
||||
Name: name,
|
||||
Size: int64(len(text)),
|
||||
ContentType: "text/plain",
|
||||
Kind: KindText,
|
||||
Text: text,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Send offers the payloads to the peer at addr and uploads them once accepted.
|
||||
func (c *Client) Send(ctx context.Context, addr netip.AddrPort, payloads []Payload, progress ProgressFunc) (OfferID, error) {
|
||||
id, decision, err := c.Offer(ctx, addr, payloads)
|
||||
if err != nil {
|
||||
return id, err
|
||||
}
|
||||
|
||||
decision, err = c.AwaitDecision(ctx, addr, id, decision)
|
||||
if err != nil {
|
||||
return id, err
|
||||
}
|
||||
if err := decisionError(decision); err != nil {
|
||||
return id, err
|
||||
}
|
||||
|
||||
return id, c.Upload(ctx, addr, id, payloads, progress)
|
||||
}
|
||||
|
||||
// Offer announces the payloads and returns the offer ID with its initial decision.
|
||||
func (c *Client) Offer(ctx context.Context, addr netip.AddrPort, payloads []Payload) (OfferID, Decision, error) {
|
||||
if len(payloads) == 0 {
|
||||
return "", DecisionPending, fmt.Errorf("%w: no payloads", ErrInvalidOffer)
|
||||
}
|
||||
return c.postOffer(ctx, baseURL(addr), payloads)
|
||||
}
|
||||
|
||||
// AwaitDecision resolves a pending decision by long-polling the receiver.
|
||||
func (c *Client) AwaitDecision(ctx context.Context, addr netip.AddrPort, id OfferID, decision Decision) (Decision, error) {
|
||||
if decision != DecisionPending {
|
||||
return decision, nil
|
||||
}
|
||||
return c.awaitDecision(ctx, baseURL(addr), id)
|
||||
}
|
||||
|
||||
// Upload streams every non-inline payload of an accepted offer.
|
||||
func (c *Client) Upload(ctx context.Context, addr netip.AddrPort, id OfferID, payloads []Payload, progress ProgressFunc) error {
|
||||
base := baseURL(addr)
|
||||
for i, p := range payloads {
|
||||
if p.Meta.Kind == KindText {
|
||||
continue
|
||||
}
|
||||
if err := c.uploadFile(ctx, base, id, i, p, progress); err != nil {
|
||||
return fmt.Errorf("upload %s: %w", p.Meta.Name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cancel withdraws an offer, taking the receiver's consent prompt with it.
|
||||
func (c *Client) Cancel(ctx context.Context, addr netip.AddrPort, id OfferID) error {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, offerURL(baseURL(addr), id), nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("build cancel request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("send cancel: %w", err)
|
||||
}
|
||||
defer drainAndClose(resp)
|
||||
|
||||
if resp.StatusCode != http.StatusNoContent && resp.StatusCode != http.StatusNotFound {
|
||||
return statusError(resp)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) postOffer(ctx context.Context, base string, payloads []Payload) (OfferID, Decision, error) {
|
||||
files := make([]FileMeta, len(payloads))
|
||||
for i, p := range payloads {
|
||||
files[i] = p.Meta
|
||||
}
|
||||
|
||||
body, err := json.Marshal(OfferRequest{SenderName: c.senderName, Files: files})
|
||||
if err != nil {
|
||||
return "", DecisionPending, fmt.Errorf("encode offer: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, base+pathOffers, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return "", DecisionPending, fmt.Errorf("build offer request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
return "", DecisionPending, fmt.Errorf("send offer: %w", err)
|
||||
}
|
||||
defer drainAndClose(resp)
|
||||
|
||||
switch resp.StatusCode {
|
||||
case http.StatusCreated, http.StatusAccepted:
|
||||
case http.StatusForbidden:
|
||||
return "", DecisionPending, ErrRefused
|
||||
default:
|
||||
return "", DecisionPending, statusError(resp)
|
||||
}
|
||||
|
||||
var offer OfferResponse
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxOfferBodySize)).Decode(&offer); err != nil {
|
||||
return "", DecisionPending, fmt.Errorf("decode offer response: %w", err)
|
||||
}
|
||||
if offer.ID == "" {
|
||||
return "", DecisionPending, fmt.Errorf("%w: receiver returned no offer id", ErrInvalidOffer)
|
||||
}
|
||||
if !offer.Decision.valid() {
|
||||
return "", DecisionPending, fmt.Errorf("%w: receiver returned decision %s", ErrInvalidOffer, offer.Decision)
|
||||
}
|
||||
|
||||
return offer.ID, offer.Decision, nil
|
||||
}
|
||||
|
||||
func (c *Client) awaitDecision(ctx context.Context, base string, id OfferID) (Decision, error) {
|
||||
deadline := time.Now().Add(c.offerTimeout)
|
||||
|
||||
for time.Now().Before(deadline) {
|
||||
decision, err := c.pollDecision(ctx, base, id)
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return DecisionPending, ctx.Err()
|
||||
}
|
||||
log.Debugf("poll file drop decision: %v", err)
|
||||
if !sleepCtx(ctx, uploadRetryDelay) {
|
||||
return DecisionPending, ctx.Err()
|
||||
}
|
||||
continue
|
||||
}
|
||||
if decision != DecisionPending {
|
||||
return decision, nil
|
||||
}
|
||||
}
|
||||
|
||||
return DecisionExpired, nil
|
||||
}
|
||||
|
||||
func (c *Client) pollDecision(ctx context.Context, base string, id OfferID) (Decision, error) {
|
||||
pollCtx, cancel := context.WithTimeout(ctx, c.pollTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(pollCtx, http.MethodGet, offerURL(base, id), nil)
|
||||
if err != nil {
|
||||
return DecisionPending, fmt.Errorf("build status request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
return DecisionPending, fmt.Errorf("poll status: %w", err)
|
||||
}
|
||||
defer drainAndClose(resp)
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return DecisionPending, ErrOfferNotFound
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return DecisionPending, statusError(resp)
|
||||
}
|
||||
|
||||
var offer OfferResponse
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxOfferBodySize)).Decode(&offer); err != nil {
|
||||
return DecisionPending, fmt.Errorf("decode status response: %w", err)
|
||||
}
|
||||
if !offer.Decision.valid() {
|
||||
return DecisionPending, fmt.Errorf("%w: receiver returned decision %s", ErrInvalidOffer, offer.Decision)
|
||||
}
|
||||
return offer.Decision, nil
|
||||
}
|
||||
|
||||
func (c *Client) uploadFile(ctx context.Context, base string, id OfferID, index int, p Payload, progress ProgressFunc) error {
|
||||
var lastErr error
|
||||
|
||||
for attempt := range maxUploadAttempts {
|
||||
offset := int64(0)
|
||||
if attempt > 0 {
|
||||
if !sleepCtx(ctx, uploadRetryDelay) {
|
||||
return ctx.Err()
|
||||
}
|
||||
confirmed, err := c.confirmedOffset(ctx, base, id, index)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
offset = confirmed
|
||||
}
|
||||
|
||||
if offset >= p.Meta.Size {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := c.putFile(ctx, base, id, index, p, offset, progress); err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
lastErr = err
|
||||
log.Debugf("upload attempt %d for %s: %v", attempt+1, p.Meta.Name, err)
|
||||
continue
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return lastErr
|
||||
}
|
||||
|
||||
func (c *Client) putFile(ctx context.Context, base string, id OfferID, index int, p Payload, offset int64, progress ProgressFunc) error {
|
||||
if p.Open == nil {
|
||||
return fmt.Errorf("payload %s has no reader", p.Meta.Name)
|
||||
}
|
||||
|
||||
body, err := p.Open(offset)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open payload: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := body.Close(); err != nil {
|
||||
log.Debugf("close payload reader: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
var reader io.Reader = body
|
||||
if progress != nil {
|
||||
reader = &progressReader{
|
||||
r: body,
|
||||
sent: offset,
|
||||
total: p.Meta.Size,
|
||||
report: func(sent int64) {
|
||||
progress(index, sent, p.Meta.Size)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
url := fileURL(base, id, index) + "?offset=" + strconv.FormatInt(offset, 10)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPut, url, reader)
|
||||
if err != nil {
|
||||
return fmt.Errorf("build upload request: %w", err)
|
||||
}
|
||||
req.ContentLength = p.Meta.Size - offset
|
||||
req.Header.Set("Content-Type", contentTypeOrDefault(p.Meta.ContentType))
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("send payload: %w", err)
|
||||
}
|
||||
defer drainAndClose(resp)
|
||||
|
||||
if resp.StatusCode == http.StatusForbidden {
|
||||
return ErrNotAccepted
|
||||
}
|
||||
if resp.StatusCode != http.StatusNoContent && resp.StatusCode != http.StatusOK {
|
||||
return statusError(resp)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) confirmedOffset(ctx context.Context, base string, id OfferID, index int) (int64, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodHead, fileURL(base, id, index), nil)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("build probe request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.http.Do(req)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("probe upload: %w", err)
|
||||
}
|
||||
defer drainAndClose(resp)
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return 0, statusError(resp)
|
||||
}
|
||||
|
||||
raw := resp.Header.Get(HeaderReceivedBytes)
|
||||
if raw == "" {
|
||||
return 0, nil
|
||||
}
|
||||
offset, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil || offset < 0 {
|
||||
return 0, fmt.Errorf("invalid %s header %q", HeaderReceivedBytes, raw)
|
||||
}
|
||||
return offset, nil
|
||||
}
|
||||
|
||||
func baseURL(addr netip.AddrPort) string {
|
||||
return "http://" + net.JoinHostPort(addr.Addr().Unmap().String(), strconv.Itoa(int(addr.Port())))
|
||||
}
|
||||
|
||||
func offerURL(base string, id OfferID) string {
|
||||
return base + pathOffersSlash + string(id)
|
||||
}
|
||||
|
||||
func fileURL(base string, id OfferID, index int) string {
|
||||
return offerURL(base, id) + "/" + segmentFiles + "/" + strconv.Itoa(index)
|
||||
}
|
||||
|
||||
func contentTypeOrDefault(ct string) string {
|
||||
if ct == "" {
|
||||
return "application/octet-stream"
|
||||
}
|
||||
return ct
|
||||
}
|
||||
|
||||
func statusError(resp *http.Response) error {
|
||||
return fmt.Errorf("receiver returned %s", resp.Status)
|
||||
}
|
||||
|
||||
func drainAndClose(resp *http.Response) {
|
||||
if _, err := io.Copy(io.Discard, io.LimitReader(resp.Body, maxOfferBodySize)); err != nil {
|
||||
log.Tracef("drain response body: %v", err)
|
||||
}
|
||||
if err := resp.Body.Close(); err != nil {
|
||||
log.Debugf("close response body: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func sleepCtx(ctx context.Context, d time.Duration) bool {
|
||||
timer := time.NewTimer(d)
|
||||
defer timer.Stop()
|
||||
|
||||
select {
|
||||
case <-timer.C:
|
||||
return true
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func decisionError(decision Decision) error {
|
||||
switch decision {
|
||||
case DecisionAccepted:
|
||||
return nil
|
||||
case DecisionDeclined:
|
||||
return ErrDeclined
|
||||
case DecisionExpired:
|
||||
return ErrExpired
|
||||
default:
|
||||
return fmt.Errorf("unexpected decision %s", decision)
|
||||
}
|
||||
}
|
||||
100
client/internal/filedrop/delivery.go
Normal file
100
client/internal/filedrop/delivery.go
Normal file
@@ -0,0 +1,100 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
func deliver(spool *Spool, offer Offer, destDir string) ([]string, error) {
|
||||
if destDir == "" {
|
||||
return nil, fmt.Errorf("no destination directory configured")
|
||||
}
|
||||
if err := os.MkdirAll(destDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("create destination dir: %w", err)
|
||||
}
|
||||
|
||||
var delivered []string
|
||||
for i, f := range offer.Files {
|
||||
if f.Kind == KindText {
|
||||
continue
|
||||
}
|
||||
|
||||
dest, err := moveToUniqueName(spool.Path(offer.ID, i), destDir, sanitizeFileName(f.Name, i))
|
||||
if err != nil {
|
||||
return delivered, fmt.Errorf("deliver %s: %w", f.Name, err)
|
||||
}
|
||||
if err := chownToDirOwner(dest, destDir); err != nil {
|
||||
log.Debugf("failed to adopt owner for %s: %v", dest, err)
|
||||
}
|
||||
delivered = append(delivered, dest)
|
||||
}
|
||||
|
||||
spool.Remove(offer.ID)
|
||||
return delivered, nil
|
||||
}
|
||||
|
||||
func sanitizeFileName(name string, index int) string {
|
||||
name = filepath.Base(filepath.Clean(strings.ReplaceAll(name, "\\", "/")))
|
||||
if name == "" || name == "." || name == ".." || name == string(filepath.Separator) {
|
||||
return fmt.Sprintf("file-%d", index)
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func moveToUniqueName(src, dir, name string) (string, error) {
|
||||
ext := filepath.Ext(name)
|
||||
stem := strings.TrimSuffix(name, ext)
|
||||
|
||||
for attempt := 0; attempt < 1000; attempt++ {
|
||||
candidate := name
|
||||
if attempt > 0 {
|
||||
candidate = fmt.Sprintf("%s (%d)%s", stem, attempt, ext)
|
||||
}
|
||||
dest := filepath.Join(dir, candidate)
|
||||
|
||||
f, err := os.OpenFile(dest, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o644)
|
||||
if err != nil {
|
||||
if os.IsExist(err) {
|
||||
continue
|
||||
}
|
||||
return "", fmt.Errorf("create destination: %w", err)
|
||||
}
|
||||
|
||||
if err := moveInto(f, src); err != nil {
|
||||
_ = f.Close()
|
||||
_ = os.Remove(dest)
|
||||
return "", err
|
||||
}
|
||||
if err := f.Close(); err != nil {
|
||||
return "", fmt.Errorf("close destination: %w", err)
|
||||
}
|
||||
if err := os.Remove(src); err != nil {
|
||||
log.Debugf("failed to remove spooled source %s: %v", src, err)
|
||||
}
|
||||
return dest, nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("no free name for %s in %s", name, dir)
|
||||
}
|
||||
|
||||
func moveInto(dst *os.File, src string) error {
|
||||
s, err := os.Open(src)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open spooled file: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := s.Close(); err != nil {
|
||||
log.Debugf("close spooled file: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := io.Copy(dst, s); err != nil {
|
||||
return fmt.Errorf("copy payload: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
7
client/internal/filedrop/delivery_other.go
Normal file
7
client/internal/filedrop/delivery_other.go
Normal file
@@ -0,0 +1,7 @@
|
||||
//go:build windows || js
|
||||
|
||||
package filedrop
|
||||
|
||||
func chownToDirOwner(string, string) error {
|
||||
return nil
|
||||
}
|
||||
29
client/internal/filedrop/delivery_unix.go
Normal file
29
client/internal/filedrop/delivery_unix.go
Normal file
@@ -0,0 +1,29 @@
|
||||
//go:build !windows && !js
|
||||
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func chownToDirOwner(path, dir string) error {
|
||||
info, err := os.Stat(dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("stat destination dir: %w", err)
|
||||
}
|
||||
|
||||
stat, ok := info.Sys().(*syscall.Stat_t)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if os.Geteuid() != 0 || int(stat.Uid) == os.Geteuid() {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := os.Chown(path, int(stat.Uid), int(stat.Gid)); err != nil {
|
||||
return fmt.Errorf("chown delivered file: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
850
client/internal/filedrop/filedrop_test.go
Normal file
850
client/internal/filedrop/filedrop_test.go
Normal file
@@ -0,0 +1,850 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
)
|
||||
|
||||
const (
|
||||
testPeer = PeerKey("peer-pubkey")
|
||||
testProfile = profilemanager.ID("test-profile")
|
||||
)
|
||||
|
||||
type staticResolver struct {
|
||||
key PeerKey
|
||||
name string
|
||||
unknown bool
|
||||
}
|
||||
|
||||
func (r staticResolver) ResolvePeer(netip.Addr) (PeerKey, string, bool) {
|
||||
if r.unknown {
|
||||
return "", "", false
|
||||
}
|
||||
return r.key, r.name, true
|
||||
}
|
||||
|
||||
type recordingNotifier struct {
|
||||
mu sync.Mutex
|
||||
offers []Offer
|
||||
completed []Offer
|
||||
failed []Offer
|
||||
withdrawn []Offer
|
||||
progress int
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) OnOffer(o Offer) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.offers = append(n.offers, o)
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) OnProgress(Offer, int, int64) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.progress++
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) OnCompleted(o Offer) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.completed = append(n.completed, o)
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) OnFailed(o Offer, _ error) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.failed = append(n.failed, o)
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) OnWithdrawn(o Offer) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.withdrawn = append(n.withdrawn, o)
|
||||
}
|
||||
|
||||
func (n *recordingNotifier) snapshot() (offers, completed, failed, withdrawn []Offer) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return append([]Offer(nil), n.offers...), append([]Offer(nil), n.completed...),
|
||||
append([]Offer(nil), n.failed...), append([]Offer(nil), n.withdrawn...)
|
||||
}
|
||||
|
||||
func startTestServer(t *testing.T, mode Mode, resolver PeerResolver) (*Server, *Client, *recordingNotifier) {
|
||||
t.Helper()
|
||||
|
||||
policy := NewPolicyStore(testProfile)
|
||||
require.NoError(t, policy.Set(Policy{Mode: mode}))
|
||||
|
||||
notifier := &recordingNotifier{}
|
||||
srv, err := NewServer(ServerConfig{
|
||||
SpoolDir: t.TempDir(),
|
||||
Policy: policy,
|
||||
Resolver: resolver,
|
||||
Notifier: notifier,
|
||||
OfferTTL: 5 * time.Second,
|
||||
})
|
||||
require.NoError(t, err, "server setup must succeed")
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
|
||||
require.NoError(t, srv.Start(ctx, netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 0)))
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, srv.Stop())
|
||||
})
|
||||
|
||||
srv.mu.RLock()
|
||||
addr := srv.listener.Addr().String()
|
||||
srv.mu.RUnlock()
|
||||
|
||||
client, err := NewClient(ClientConfig{
|
||||
SenderName: "sender",
|
||||
PollTimeout: 2 * time.Second,
|
||||
OfferTimeout: 5 * time.Second,
|
||||
Dial: func(ctx context.Context, network, _ string) (net.Conn, error) {
|
||||
var d net.Dialer
|
||||
return d.DialContext(ctx, network, addr)
|
||||
},
|
||||
})
|
||||
require.NoError(t, err, "client setup must succeed")
|
||||
|
||||
return srv, client, notifier
|
||||
}
|
||||
|
||||
func filePayload(t *testing.T, name string, content []byte) Payload {
|
||||
t.Helper()
|
||||
|
||||
path := filepath.Join(t.TempDir(), name)
|
||||
require.NoError(t, os.WriteFile(path, content, 0o600))
|
||||
|
||||
return Payload{
|
||||
Meta: FileMeta{Name: name, Size: int64(len(content)), ContentType: "application/octet-stream"},
|
||||
Open: func(offset int64) (io.ReadCloser, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := f.Seek(offset, io.SeekStart); err != nil {
|
||||
_ = f.Close()
|
||||
return nil, err
|
||||
}
|
||||
return f, nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
var testAddr = netip.AddrPortFrom(netip.AddrFrom4([4]byte{100, 64, 0, 1}), Port)
|
||||
|
||||
func TestAutoAcceptTransfersPayload(t *testing.T) {
|
||||
srv, client, notifier := startTestServer(t, ModeAutoAccept, staticResolver{key: testPeer, name: "laptop"})
|
||||
|
||||
content := []byte(strings.Repeat("netbird", 1000))
|
||||
payload := filePayload(t, "report.bin", content)
|
||||
|
||||
var lastSent int64
|
||||
id, err := client.Send(context.Background(), testAddr, []Payload{payload}, func(_ int, sent, _ int64) {
|
||||
lastSent = sent
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, id, "receiver must return an offer id")
|
||||
|
||||
assert.Equal(t, int64(len(content)), lastSent, "progress must reach the full payload size")
|
||||
|
||||
staged, err := os.ReadFile(srv.Spool().Path(id, 0))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, content, staged, "staged bytes should match what was sent")
|
||||
|
||||
offer, ok := srv.Offers().Get(testPeer, id)
|
||||
require.True(t, ok, "offer must still be tracked")
|
||||
assert.Equal(t, StateCompleted, offer.State, "offer should be completed")
|
||||
assert.Equal(t, "laptop", offer.SenderName, "sender name should come from the resolver")
|
||||
|
||||
_, completed, _, _ := notifier.snapshot()
|
||||
require.Len(t, completed, 1, "one completion event expected")
|
||||
assert.Equal(t, id, completed[0].ID)
|
||||
}
|
||||
|
||||
func TestAskModeAcceptReleasesUpload(t *testing.T) {
|
||||
srv, client, notifier := startTestServer(t, ModeAsk, staticResolver{key: testPeer})
|
||||
|
||||
content := []byte("consent required")
|
||||
payload := filePayload(t, "note.txt", content)
|
||||
|
||||
go func() {
|
||||
for {
|
||||
offers, _, _, _ := notifier.snapshot()
|
||||
if len(offers) > 0 {
|
||||
srv.Offers().Decide(offers[0].ID, DecisionAccepted)
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
}()
|
||||
|
||||
id, err := client.Send(context.Background(), testAddr, []Payload{payload}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
staged, err := os.ReadFile(srv.Spool().Path(id, 0))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, content, staged, "payload should arrive after acceptance")
|
||||
|
||||
offers, _, _, _ := notifier.snapshot()
|
||||
require.Len(t, offers, 1, "the pending offer must be raised exactly once")
|
||||
assert.Equal(t, DecisionPending, offers[0].Decision, "the raised offer starts pending")
|
||||
}
|
||||
|
||||
func TestAskModeDeclineKeepsPayloadOut(t *testing.T) {
|
||||
srv, client, notifier := startTestServer(t, ModeAsk, staticResolver{key: testPeer})
|
||||
|
||||
go func() {
|
||||
for {
|
||||
offers, _, _, _ := notifier.snapshot()
|
||||
if len(offers) > 0 {
|
||||
srv.Offers().Decide(offers[0].ID, DecisionDeclined)
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
}()
|
||||
|
||||
id, err := client.Send(context.Background(), testAddr, []Payload{filePayload(t, "x.bin", []byte("data"))}, nil)
|
||||
require.ErrorIs(t, err, ErrDeclined, "sender must see the decline")
|
||||
|
||||
_, statErr := os.Stat(srv.Spool().Path(id, 0))
|
||||
assert.True(t, os.IsNotExist(statErr), "declined payload must never be staged")
|
||||
}
|
||||
|
||||
func TestOffModeRefusesOffer(t *testing.T) {
|
||||
_, client, notifier := startTestServer(t, ModeOff, staticResolver{key: testPeer})
|
||||
|
||||
_, err := client.Send(context.Background(), testAddr, []Payload{filePayload(t, "x.bin", []byte("data"))}, nil)
|
||||
require.ErrorIs(t, err, ErrRefused, "an off receiver must refuse the offer")
|
||||
|
||||
offers, _, _, _ := notifier.snapshot()
|
||||
assert.Empty(t, offers, "a refused offer must not reach the user")
|
||||
}
|
||||
|
||||
func TestUnknownSenderIsRefused(t *testing.T) {
|
||||
_, client, _ := startTestServer(t, ModeAutoAccept, staticResolver{unknown: true})
|
||||
|
||||
_, err := client.Send(context.Background(), testAddr, []Payload{filePayload(t, "x.bin", []byte("data"))}, nil)
|
||||
require.ErrorIs(t, err, ErrRefused, "an unresolvable source address must be refused")
|
||||
}
|
||||
|
||||
func TestUploadResumesFromConfirmedOffset(t *testing.T) {
|
||||
srv, client, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: testPeer})
|
||||
|
||||
content := []byte(strings.Repeat("resume", 500))
|
||||
payload := filePayload(t, "big.bin", content)
|
||||
|
||||
offer := srv.Offers().Add(testPeer, "", []FileMeta{payload.Meta}, DecisionAccepted)
|
||||
require.NoError(t, srv.Spool().Prepare(offer.ID))
|
||||
|
||||
half := int64(len(content) / 2)
|
||||
_, err := srv.Spool().Write(offer.ID, 0, 0, strings.NewReader(string(content[:half])), half)
|
||||
require.NoError(t, err)
|
||||
|
||||
srv.mu.RLock()
|
||||
base := "http://" + srv.listener.Addr().String()
|
||||
srv.mu.RUnlock()
|
||||
|
||||
confirmed, err := client.confirmedOffset(context.Background(), base, offer.ID, 0)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, half, confirmed, "receiver must report the staged prefix")
|
||||
|
||||
require.NoError(t, client.putFile(context.Background(), base, offer.ID, 0, payload, confirmed, nil))
|
||||
|
||||
staged, err := os.ReadFile(srv.Spool().Path(offer.ID, 0))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, content, staged, "resumed upload must reconstruct the full payload")
|
||||
}
|
||||
|
||||
func TestUploadIsBoundedByAnnouncedSize(t *testing.T) {
|
||||
srv, _, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: testPeer})
|
||||
|
||||
// The request is issued raw: the stdlib client refuses to send a body that
|
||||
announced := int64(10)
|
||||
offer := srv.Offers().Add(testPeer, "", []FileMeta{{Name: "lie.bin", Size: announced}}, DecisionAccepted)
|
||||
require.NoError(t, srv.Spool().Prepare(offer.ID))
|
||||
|
||||
srv.mu.RLock()
|
||||
addr := srv.listener.Addr().String()
|
||||
srv.mu.RUnlock()
|
||||
|
||||
conn, err := net.Dial("tcp", addr)
|
||||
require.NoError(t, err)
|
||||
defer func() {
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
oversized := strings.Repeat("A", 100)
|
||||
request := "PUT /v1/offers/" + string(offer.ID) + "/files/0?offset=0 HTTP/1.1\r\n" +
|
||||
"Host: filedrop\r\nContent-Length: 100\r\nConnection: close\r\n\r\n" + oversized
|
||||
_, err = conn.Write([]byte(request))
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = io.ReadAll(conn)
|
||||
require.NoError(t, err)
|
||||
|
||||
staged, err := os.ReadFile(srv.Spool().Path(offer.ID, 0))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, staged, int(announced), "staged size must be capped at the announced size")
|
||||
}
|
||||
|
||||
func TestCancelWithdrawsPendingOffer(t *testing.T) {
|
||||
srv, client, notifier := startTestServer(t, ModeAsk, staticResolver{key: testPeer})
|
||||
|
||||
sendCtx, cancelSend := context.WithCancel(context.Background())
|
||||
defer cancelSend()
|
||||
|
||||
go func() {
|
||||
_, _ = client.Send(sendCtx, testAddr, []Payload{filePayload(t, "x.bin", []byte("data"))}, nil)
|
||||
}()
|
||||
|
||||
var id OfferID
|
||||
require.Eventually(t, func() bool {
|
||||
offers, _, _, _ := notifier.snapshot()
|
||||
if len(offers) == 0 {
|
||||
return false
|
||||
}
|
||||
id = offers[0].ID
|
||||
return true
|
||||
}, 3*time.Second, 10*time.Millisecond, "offer must reach the receiver")
|
||||
|
||||
require.NoError(t, client.Cancel(context.Background(), testAddr, id))
|
||||
|
||||
_, ok := srv.Offers().Get(testPeer, id)
|
||||
assert.False(t, ok, "a withdrawn offer must be dropped")
|
||||
|
||||
_, _, _, withdrawn := notifier.snapshot()
|
||||
require.Len(t, withdrawn, 1, "the consent prompt must be withdrawn")
|
||||
assert.Equal(t, id, withdrawn[0].ID)
|
||||
}
|
||||
|
||||
func TestTextPayloadStaysInline(t *testing.T) {
|
||||
srv, client, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: testPeer})
|
||||
|
||||
id, err := client.Send(context.Background(), testAddr, []Payload{TextPayload("snippet", "hello peer")}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
offer, ok := srv.Offers().Get(testPeer, id)
|
||||
require.True(t, ok)
|
||||
require.Len(t, offer.Files, 1)
|
||||
assert.Equal(t, "hello peer", offer.Files[0].Text, "text must arrive in the offer itself")
|
||||
assert.Equal(t, StateCompleted, offer.State, "a text-only offer completes without an upload")
|
||||
|
||||
_, statErr := os.Stat(srv.Spool().Path(id, 0))
|
||||
assert.True(t, os.IsNotExist(statErr), "text payloads must not be written to the spool")
|
||||
}
|
||||
|
||||
func TestOfferExpiresWithoutDecision(t *testing.T) {
|
||||
policy := NewPolicyStore(testProfile)
|
||||
require.NoError(t, policy.SetMode(ModeAsk))
|
||||
|
||||
srv, err := NewServer(ServerConfig{
|
||||
SpoolDir: t.TempDir(),
|
||||
Policy: policy,
|
||||
Resolver: staticResolver{key: testPeer},
|
||||
OfferTTL: 100 * time.Millisecond,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
offer := srv.Offers().Add(testPeer, "", []FileMeta{{Name: "x", Size: 1}}, DecisionPending)
|
||||
|
||||
awaited, err := srv.Offers().Await(context.Background(), testPeer, offer.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, DecisionExpired, awaited.Decision, "an unanswered offer must expire")
|
||||
assert.Equal(t, StateExpired, awaited.State)
|
||||
}
|
||||
|
||||
func TestDecideIsFinal(t *testing.T) {
|
||||
store := NewOfferStore(time.Minute)
|
||||
offer := store.Add(testPeer, "", []FileMeta{{Name: "x", Size: 1}}, DecisionPending)
|
||||
|
||||
_, ok := store.Decide(offer.ID, DecisionDeclined)
|
||||
require.True(t, ok, "the first decision must be recorded")
|
||||
|
||||
_, ok = store.Decide(offer.ID, DecisionAccepted)
|
||||
assert.False(t, ok, "a decided offer must not be revived")
|
||||
|
||||
current, ok := store.Get(testPeer, offer.ID)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, DecisionDeclined, current.Decision, "the original decision must stand")
|
||||
}
|
||||
|
||||
func TestOfferIsScopedToItsSender(t *testing.T) {
|
||||
store := NewOfferStore(time.Minute)
|
||||
offer := store.Add(testPeer, "", []FileMeta{{Name: "x", Size: 1}}, DecisionAccepted)
|
||||
|
||||
_, ok := store.Get("other-peer", offer.ID)
|
||||
assert.False(t, ok, "another peer must not see the offer")
|
||||
|
||||
_, err := store.Await(context.Background(), "other-peer", offer.ID)
|
||||
assert.ErrorIs(t, err, ErrOfferNotFound, "another peer must not poll the offer")
|
||||
}
|
||||
|
||||
func TestPolicyEvaluation(t *testing.T) {
|
||||
store := NewPolicyStore(testProfile)
|
||||
require.NoError(t, store.Set(Policy{Mode: ModeAsk}))
|
||||
|
||||
assert.Equal(t, ModeAsk, store.Evaluate(testPeer), "unknown senders are asked about")
|
||||
|
||||
require.NoError(t, store.SetSenderRule(testPeer, SenderRuleBlock))
|
||||
assert.Equal(t, ModeOff, store.Evaluate(testPeer), "a blocked sender is refused")
|
||||
|
||||
require.NoError(t, store.SetSenderRule(testPeer, SenderRuleAlwaysAccept))
|
||||
assert.Equal(t, ModeAutoAccept, store.Evaluate(testPeer), "an always-accept sender skips the prompt")
|
||||
|
||||
require.NoError(t, store.SetSenderRule(testPeer, SenderRuleDefault))
|
||||
assert.Equal(t, ModeAsk, store.Evaluate(testPeer), "clearing the rule restores the base mode")
|
||||
}
|
||||
|
||||
func TestPolicyRejectsUnknownModeAndDeniesOnCorruptRule(t *testing.T) {
|
||||
store := NewPolicyStore(testProfile)
|
||||
|
||||
require.Error(t, store.SetMode(Mode(200)), "an unknown mode must be rejected")
|
||||
assert.Equal(t, ModeAsk, store.Get().Mode, "the rejected mode must not be applied")
|
||||
|
||||
require.NoError(t, store.SetSenderRule(testPeer, SenderRule(200)))
|
||||
assert.Equal(t, ModeOff, store.Evaluate(testPeer), "an unrecognized rule must deny")
|
||||
}
|
||||
|
||||
type memStore struct {
|
||||
mu sync.Mutex
|
||||
sections map[string][]byte
|
||||
loadErr error
|
||||
}
|
||||
|
||||
func newMemStore() *memStore {
|
||||
return &memStore{sections: map[string][]byte{}}
|
||||
}
|
||||
|
||||
func (s *memStore) Get(namespace string, v any) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.loadErr != nil {
|
||||
return false, s.loadErr
|
||||
}
|
||||
raw, ok := s.sections[namespace]
|
||||
if !ok {
|
||||
return false, nil
|
||||
}
|
||||
return true, json.Unmarshal(raw, v)
|
||||
}
|
||||
|
||||
func (s *memStore) Put(namespace string, v any) error {
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.sections[namespace] = raw
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestPolicyIsScopedPerProfile(t *testing.T) {
|
||||
work, home := profilemanager.ID("work"), profilemanager.ID("home")
|
||||
workPrefs, homePrefs := newMemStore(), newMemStore()
|
||||
|
||||
workStore := LoadPolicyStore(work, workPrefs)
|
||||
require.NoError(t, workStore.SetMode(ModeOff))
|
||||
require.NoError(t, workStore.SetSenderRule(testPeer, SenderRuleBlock))
|
||||
|
||||
homeStore := LoadPolicyStore(home, homePrefs)
|
||||
assert.Equal(t, ModeAsk, homeStore.Get().Mode, "another profile keeps the default mode")
|
||||
assert.Equal(t, ModeAsk, homeStore.Evaluate(testPeer), "a block in one profile must not apply to another")
|
||||
|
||||
reloaded := LoadPolicyStore(work, workPrefs)
|
||||
assert.Equal(t, ModeOff, reloaded.Get().Mode, "the profile's mode must survive a reload")
|
||||
assert.Equal(t, ModeOff, reloaded.Evaluate(testPeer), "the profile's sender rule must survive a reload")
|
||||
}
|
||||
|
||||
func TestPolicyFallsBackToDefaultsOnLoadFailure(t *testing.T) {
|
||||
prefs := newMemStore()
|
||||
prefs.loadErr = errors.New("store unavailable")
|
||||
|
||||
store := LoadPolicyStore(testProfile, prefs)
|
||||
|
||||
assert.Equal(t, ModeAsk, store.Get().Mode, "an unreadable policy must not open the device up")
|
||||
assert.Equal(t, ModeAsk, store.Evaluate(testPeer), "the safe default applies to unknown senders")
|
||||
}
|
||||
|
||||
func TestPolicyRejectsStoredInvalidModeOnLoad(t *testing.T) {
|
||||
prefs := newMemStore()
|
||||
require.NoError(t, prefs.Put(namespacePolicy, Policy{Mode: Mode(200)}))
|
||||
|
||||
store := LoadPolicyStore(testProfile, prefs)
|
||||
|
||||
assert.Equal(t, ModeAsk, store.Get().Mode, "a corrupted stored mode must fall back to the default")
|
||||
}
|
||||
|
||||
func TestStoreKeepsPolicyAndHistoryApart(t *testing.T) {
|
||||
prefs := newMemStore()
|
||||
|
||||
mgr, err := NewManager(ManagerConfig{Profile: testProfile, DataDir: t.TempDir(), Store: prefs})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, mgr.Policy().SetMode(ModeAutoAccept))
|
||||
require.NoError(t, mgr.SetDestinationDir("/tmp/received"))
|
||||
mgr.history.Upsert(Transfer{ID: "offer-1", PeerKey: testPeer, State: StateCompleted})
|
||||
require.NoError(t, mgr.Close())
|
||||
|
||||
reloaded, err := NewManager(ManagerConfig{Profile: testProfile, DataDir: t.TempDir(), Store: prefs})
|
||||
require.NoError(t, err)
|
||||
defer func() { require.NoError(t, reloaded.Close()) }()
|
||||
|
||||
assert.Equal(t, ModeAutoAccept, reloaded.Policy().Get().Mode, "the policy must survive a reload")
|
||||
assert.Equal(t, "/tmp/received", reloaded.DestinationDir(), "the destination must survive a reload")
|
||||
require.Len(t, reloaded.Transfers(), 1, "the history must survive a reload")
|
||||
assert.Equal(t, OfferID("offer-1"), reloaded.Transfers()[0].ID)
|
||||
}
|
||||
|
||||
func TestHistoryDropsOldestTerminalEntriesOverCap(t *testing.T) {
|
||||
history := LoadHistory(newMemStore())
|
||||
|
||||
for i := 0; i < historyCap+10; i++ {
|
||||
history.Upsert(Transfer{ID: OfferID(fmt.Sprintf("offer-%d", i)), State: StateCompleted})
|
||||
}
|
||||
|
||||
entries := history.List()
|
||||
require.Len(t, entries, historyCap, "the log must stay bounded")
|
||||
assert.Equal(t, OfferID(fmt.Sprintf("offer-%d", historyCap+9)), entries[0].ID, "the newest entry stays")
|
||||
}
|
||||
|
||||
func TestHistoryKeepsLiveTransfersOverCap(t *testing.T) {
|
||||
history := LoadHistory(newMemStore())
|
||||
|
||||
history.Upsert(Transfer{ID: "live", State: StateTransferring})
|
||||
for i := 0; i < historyCap+5; i++ {
|
||||
history.Upsert(Transfer{ID: OfferID(fmt.Sprintf("done-%d", i)), State: StateCompleted})
|
||||
}
|
||||
|
||||
_, ok := history.Get("live")
|
||||
assert.True(t, ok, "a transfer still running must not be pruned")
|
||||
}
|
||||
|
||||
func TestHistorySettlesTransfersInterruptedByRestart(t *testing.T) {
|
||||
store := newMemStore()
|
||||
|
||||
history := LoadHistory(store)
|
||||
history.Upsert(Transfer{ID: "pending", State: StatePending})
|
||||
history.Upsert(Transfer{ID: "moving", State: StateTransferring})
|
||||
history.Upsert(Transfer{ID: "done", State: StateCompleted})
|
||||
history.Upsert(Transfer{ID: "refused", State: StateDeclined})
|
||||
|
||||
// A fresh load stands in for the next process: nothing survives to finish
|
||||
// whatever was still moving.
|
||||
reloaded := LoadHistory(store)
|
||||
|
||||
for _, tc := range []struct {
|
||||
id OfferID
|
||||
state State
|
||||
reason FailureReason
|
||||
}{
|
||||
{"pending", StateFailed, ReasonInterrupted},
|
||||
{"moving", StateFailed, ReasonInterrupted},
|
||||
{"done", StateCompleted, ReasonNone},
|
||||
{"refused", StateDeclined, ReasonNone},
|
||||
} {
|
||||
entry, ok := reloaded.Get(tc.id)
|
||||
require.True(t, ok, "entry %s must survive the reload", tc.id)
|
||||
assert.Equal(t, tc.state, entry.State, "state of %s", tc.id)
|
||||
assert.Equal(t, tc.reason, entry.Reason, "reason of %s", tc.id)
|
||||
}
|
||||
|
||||
// The settled states are written back, so a third start sees them as final
|
||||
// rather than settling them again.
|
||||
third := LoadHistory(store)
|
||||
entry, ok := third.Get("moving")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, StateFailed, entry.State)
|
||||
}
|
||||
|
||||
func TestSpoolWriteTruncatesStaleTail(t *testing.T) {
|
||||
spool, err := NewSpool(t.TempDir())
|
||||
require.NoError(t, err)
|
||||
|
||||
id := OfferID("offer")
|
||||
require.NoError(t, spool.Prepare(id))
|
||||
|
||||
_, err = spool.Write(id, 0, 0, strings.NewReader("AAAAAAAAAA"), 10)
|
||||
require.NoError(t, err)
|
||||
|
||||
total, err := spool.Write(id, 0, 2, strings.NewReader("BB"), 10)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(4), total, "staged size follows the resumed write")
|
||||
|
||||
staged, err := os.ReadFile(spool.Path(id, 0))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "AABB", string(staged), "stale bytes past the offset must be dropped")
|
||||
}
|
||||
|
||||
func TestSpoolCleanupDropsStalePartials(t *testing.T) {
|
||||
spool, err := NewSpool(t.TempDir())
|
||||
require.NoError(t, err)
|
||||
|
||||
stale, fresh := OfferID("stale"), OfferID("fresh")
|
||||
require.NoError(t, spool.Prepare(stale))
|
||||
require.NoError(t, spool.Prepare(fresh))
|
||||
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
require.NoError(t, os.Chtimes(spool.OfferDir(stale), old, old))
|
||||
|
||||
spool.Cleanup(time.Hour, time.Now())
|
||||
|
||||
_, err = os.Stat(spool.OfferDir(stale))
|
||||
assert.True(t, os.IsNotExist(err), "the stale offer dir must be removed")
|
||||
|
||||
_, err = os.Stat(spool.OfferDir(fresh))
|
||||
assert.NoError(t, err, "a recent offer dir must survive")
|
||||
}
|
||||
|
||||
func TestParseOfferPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
id OfferID
|
||||
index int
|
||||
hasIndex bool
|
||||
wantErr bool
|
||||
}{
|
||||
{path: "/v1/offers/abc", id: "abc"},
|
||||
{path: "/v1/offers/abc/files/3", id: "abc", index: 3, hasIndex: true},
|
||||
{path: "/v1/offers/", wantErr: true},
|
||||
{path: "/v1/offers/abc/files", wantErr: true},
|
||||
{path: "/v1/offers/abc/other/1", wantErr: true},
|
||||
{path: "/v1/offers/abc/files/-1", wantErr: true},
|
||||
{path: "/v1/offers/abc/files/x", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.path, func(t *testing.T) {
|
||||
id, index, hasIndex, err := parseOfferPath(tc.path)
|
||||
if tc.wantErr {
|
||||
assert.Error(t, err)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tc.id, id)
|
||||
assert.Equal(t, tc.index, index)
|
||||
assert.Equal(t, tc.hasIndex, hasIndex)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateOffer(t *testing.T) {
|
||||
assert.Error(t, validateOffer(nil), "an empty offer is invalid")
|
||||
assert.Error(t, validateOffer([]FileMeta{{Name: "x", Size: -1}}), "a negative size is invalid")
|
||||
assert.Error(t, validateOffer(make([]FileMeta, MaxOfferFiles+1)), "too many files is invalid")
|
||||
assert.Error(t, validateOffer([]FileMeta{{
|
||||
Name: "x", Kind: KindText, Text: strings.Repeat("a", MaxInlineTextSize+1),
|
||||
}}), "oversized inline text is invalid")
|
||||
|
||||
assert.NoError(t, validateOffer([]FileMeta{{Name: "x", Size: 10}}))
|
||||
}
|
||||
|
||||
func TestStopIsIdempotent(t *testing.T) {
|
||||
srv, err := NewServer(ServerConfig{
|
||||
SpoolDir: t.TempDir(),
|
||||
Policy: NewPolicyStore(testProfile),
|
||||
Resolver: staticResolver{key: testPeer},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, srv.Stop(), "stopping a server that never started is a no-op")
|
||||
|
||||
require.NoError(t, srv.Start(context.Background(), netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 0)))
|
||||
require.NoError(t, srv.Stop())
|
||||
require.NoError(t, srv.Stop(), "the second stop must also be a no-op")
|
||||
}
|
||||
|
||||
func TestStartRejectsSecondStart(t *testing.T) {
|
||||
srv, err := NewServer(ServerConfig{
|
||||
SpoolDir: t.TempDir(),
|
||||
Policy: NewPolicyStore(testProfile),
|
||||
Resolver: staticResolver{key: testPeer},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
addr := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 0)
|
||||
require.NoError(t, srv.Start(context.Background(), addr))
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, srv.Stop())
|
||||
})
|
||||
|
||||
err = srv.Start(context.Background(), addr)
|
||||
require.Error(t, err, "a running server must reject a second start")
|
||||
|
||||
srv.mu.RLock()
|
||||
running := srv.httpServer != nil && srv.listener != nil
|
||||
srv.mu.RUnlock()
|
||||
assert.True(t, running, "the original listener must survive the rejected start")
|
||||
}
|
||||
|
||||
func TestNewServerRequiresResolver(t *testing.T) {
|
||||
_, err := NewServer(ServerConfig{SpoolDir: t.TempDir(), Policy: NewPolicyStore(testProfile)})
|
||||
require.Error(t, err, "a server without peer resolution must not be constructed")
|
||||
}
|
||||
|
||||
func TestNewServerRequiresPolicy(t *testing.T) {
|
||||
_, err := NewServer(ServerConfig{SpoolDir: t.TempDir(), Resolver: staticResolver{key: testPeer}})
|
||||
require.Error(t, err, "a server without a profile policy must not be constructed")
|
||||
}
|
||||
|
||||
func TestNewClientRequiresDialer(t *testing.T) {
|
||||
_, err := NewClient(ClientConfig{})
|
||||
require.Error(t, err, "a client without a dialer must not be constructed")
|
||||
}
|
||||
|
||||
func TestSendRejectsEmptyPayloadSet(t *testing.T) {
|
||||
_, client, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: testPeer})
|
||||
|
||||
_, err := client.Send(context.Background(), testAddr, nil, nil)
|
||||
assert.True(t, errors.Is(err, ErrInvalidOffer), "sending nothing is an invalid offer")
|
||||
}
|
||||
|
||||
func TestServerFallsBackWhenPortBusy(t *testing.T) {
|
||||
blocker, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err, "blocker listener must bind")
|
||||
defer func() {
|
||||
require.NoError(t, blocker.Close())
|
||||
}()
|
||||
busyPort := uint16(blocker.Addr().(*net.TCPAddr).Port)
|
||||
|
||||
srv, err := NewServer(ServerConfig{
|
||||
SpoolDir: t.TempDir(),
|
||||
Policy: NewPolicyStore(testProfile),
|
||||
Resolver: staticResolver{key: testPeer},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
addr := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), busyPort)
|
||||
require.NoError(t, srv.Start(context.Background(), addr), "start must fall back instead of failing")
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, srv.Stop())
|
||||
})
|
||||
|
||||
bound := srv.BoundPort()
|
||||
assert.NotZero(t, bound, "fallback must report the bound port")
|
||||
assert.NotEqual(t, busyPort, bound, "fallback must pick a different port")
|
||||
}
|
||||
|
||||
func TestPortRegistryAwait(t *testing.T) {
|
||||
reg := NewPortRegistry()
|
||||
|
||||
reg.Set(testPeer, 5000)
|
||||
port, changed := reg.Await(context.Background(), testPeer, 0)
|
||||
assert.True(t, changed, "known differing port must return immediately")
|
||||
assert.Equal(t, uint16(5000), port)
|
||||
|
||||
go func() {
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
reg.Set(testPeer, 5000)
|
||||
}()
|
||||
_, changed = reg.Await(context.Background(), testPeer, 5000)
|
||||
assert.False(t, changed, "an advertisement equal to the used port must release the waiter as unchanged")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||
defer cancel()
|
||||
_, changed = reg.Await(ctx, testPeer, 5000)
|
||||
assert.False(t, changed, "timeout without advertisement must report unchanged")
|
||||
}
|
||||
|
||||
// senderManager builds a send-only manager whose dialer reaches the test server only
|
||||
// on realPort; other ports behave per defaultPortBehavior ("refuse" or "hang").
|
||||
func senderManager(t *testing.T, serverAddr string, realPort uint16, defaultPortBehavior string) *Manager {
|
||||
t.Helper()
|
||||
|
||||
mgr, err := NewManager(ManagerConfig{Profile: testProfile, DataDir: t.TempDir()})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, mgr.Close())
|
||||
})
|
||||
|
||||
mgr.SetTunnel(func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
ap, err := netip.ParseAddrPort(addr)
|
||||
require.NoError(t, err, "dialer must receive a valid addr")
|
||||
if ap.Port() == realPort {
|
||||
var d net.Dialer
|
||||
return d.DialContext(ctx, network, serverAddr)
|
||||
}
|
||||
if defaultPortBehavior == "hang" {
|
||||
<-ctx.Done()
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
return nil, &net.OpError{Op: "dial", Net: network, Err: errors.New("connection refused")}
|
||||
}, "sender")
|
||||
|
||||
return mgr
|
||||
}
|
||||
|
||||
func waitForState(t *testing.T, mgr *Manager, id OfferID, want State) {
|
||||
t.Helper()
|
||||
require.Eventually(t, func() bool {
|
||||
tr, ok := mgr.history.Get(id)
|
||||
return ok && tr.State == want
|
||||
}, 10*time.Second, 20*time.Millisecond, "transfer must reach state %s", want)
|
||||
}
|
||||
|
||||
func TestSendRetriesOnAdvertisedPort(t *testing.T) {
|
||||
srv, _, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: PeerKey("sender-key")})
|
||||
srv.mu.RLock()
|
||||
serverAddr := srv.listener.Addr().String()
|
||||
srv.mu.RUnlock()
|
||||
realPort := srv.BoundPort()
|
||||
|
||||
mgr := senderManager(t, serverAddr, realPort, "refuse")
|
||||
|
||||
id, err := mgr.Send(testPeer, "receiver", netip.AddrFrom4([4]byte{100, 64, 0, 9}), []Payload{TextPayload("t", "hello")})
|
||||
require.NoError(t, err)
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
mgr.Ports().Set(testPeer, realPort)
|
||||
|
||||
waitForState(t, mgr, id, StateCompleted)
|
||||
}
|
||||
|
||||
func TestSendAbortsHangingAttemptOnAdvertisedPort(t *testing.T) {
|
||||
srv, _, _ := startTestServer(t, ModeAutoAccept, staticResolver{key: PeerKey("sender-key")})
|
||||
srv.mu.RLock()
|
||||
serverAddr := srv.listener.Addr().String()
|
||||
srv.mu.RUnlock()
|
||||
realPort := srv.BoundPort()
|
||||
|
||||
mgr := senderManager(t, serverAddr, realPort, "hang")
|
||||
|
||||
id, err := mgr.Send(testPeer, "receiver", netip.AddrFrom4([4]byte{100, 64, 0, 9}), []Payload{TextPayload("t", "hello")})
|
||||
require.NoError(t, err)
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
mgr.Ports().Set(testPeer, realPort)
|
||||
|
||||
waitForState(t, mgr, id, StateCompleted)
|
||||
}
|
||||
|
||||
func TestSendFailsWhenSignalConfirmsUsedPort(t *testing.T) {
|
||||
mgr := senderManager(t, "127.0.0.1:1", 1, "refuse")
|
||||
|
||||
id, err := mgr.Send(testPeer, "receiver", netip.AddrFrom4([4]byte{100, 64, 0, 9}), []Payload{TextPayload("t", "hello")})
|
||||
require.NoError(t, err)
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
mgr.Ports().Set(testPeer, 0)
|
||||
|
||||
waitForState(t, mgr, id, StateFailed)
|
||||
}
|
||||
205
client/internal/filedrop/history.go
Normal file
205
client/internal/filedrop/history.go
Normal file
@@ -0,0 +1,205 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// The transfer directions.
|
||||
const (
|
||||
DirectionReceived Direction = iota
|
||||
DirectionSent
|
||||
)
|
||||
|
||||
// The failure reasons a transfer can end with; None accompanies every other state.
|
||||
const (
|
||||
ReasonNone FailureReason = iota
|
||||
// ReasonUnreachable marks a transport-level failure: nothing listens on the
|
||||
// peer's file drop port, so the client is old or receiving is off.
|
||||
ReasonUnreachable
|
||||
// ReasonInterrupted marks a transfer that was still moving when the process
|
||||
// stopped; nothing survived to finish or resume it.
|
||||
ReasonInterrupted
|
||||
)
|
||||
|
||||
const historyCap = 30
|
||||
|
||||
// Direction tells whether a transfer was sent by this device or received on it.
|
||||
type Direction uint8
|
||||
|
||||
// FailureReason classifies why a transfer failed, when it is known.
|
||||
type FailureReason uint8
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (d Direction) String() string {
|
||||
switch d {
|
||||
case DirectionReceived:
|
||||
return "received"
|
||||
case DirectionSent:
|
||||
return "sent"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(d))
|
||||
}
|
||||
}
|
||||
|
||||
// Transfer is one history entry: a sent or received offer with its outcome.
|
||||
type Transfer struct {
|
||||
ID OfferID `json:"id"`
|
||||
Direction Direction `json:"direction"`
|
||||
PeerKey PeerKey `json:"peerKey"`
|
||||
PeerName string `json:"peerName"`
|
||||
Files []FileMeta `json:"files"`
|
||||
State State `json:"state"`
|
||||
Transferred int64 `json:"transferred"`
|
||||
TotalSize int64 `json:"totalSize"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
DeliveredPaths []string `json:"deliveredPaths,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Reason FailureReason `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
// History is the persisted transfer log of one profile, newest first.
|
||||
type History struct {
|
||||
mu sync.RWMutex
|
||||
store Store
|
||||
entries []Transfer
|
||||
}
|
||||
|
||||
// LoadHistory reads the stored transfer log, starting empty when unreadable.
|
||||
func LoadHistory(store Store) *History {
|
||||
h := &History{store: store}
|
||||
|
||||
var entries []Transfer
|
||||
if err := loadSection(store, namespaceHistory, &entries); err != nil {
|
||||
log.Warnf("failed to read file drop history, starting empty: %v", err)
|
||||
return h
|
||||
}
|
||||
h.entries = entries
|
||||
if h.settleInterrupted() {
|
||||
h.persist()
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
func (t Transfer) terminal() bool {
|
||||
switch t.State {
|
||||
case StateCompleted, StateDeclined, StateExpired, StateCancelled, StateFailed:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (t Transfer) clone() Transfer {
|
||||
c := t
|
||||
c.Files = slices.Clone(t.Files)
|
||||
c.DeliveredPaths = slices.Clone(t.DeliveredPaths)
|
||||
return c
|
||||
}
|
||||
|
||||
// Upsert inserts or replaces the entry with the same ID and persists the log.
|
||||
func (h *History) Upsert(t Transfer) {
|
||||
t.UpdatedAt = time.Now()
|
||||
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
if i := h.indexOf(t.ID); i >= 0 {
|
||||
h.entries[i] = t
|
||||
} else {
|
||||
h.entries = slices.Insert(h.entries, 0, t)
|
||||
h.prune()
|
||||
}
|
||||
h.persist()
|
||||
}
|
||||
|
||||
// SetProgress updates the transferred byte count in memory only.
|
||||
func (h *History) SetProgress(id OfferID, transferred int64) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
if i := h.indexOf(id); i >= 0 {
|
||||
h.entries[i].Transferred = transferred
|
||||
h.entries[i].UpdatedAt = time.Now()
|
||||
if h.entries[i].State == StatePending {
|
||||
h.entries[i].State = StateTransferring
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Get returns the entry with the given ID.
|
||||
func (h *History) Get(id OfferID) (Transfer, bool) {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
|
||||
if i := h.indexOf(id); i >= 0 {
|
||||
return h.entries[i].clone(), true
|
||||
}
|
||||
return Transfer{}, false
|
||||
}
|
||||
|
||||
// List returns every entry, newest first.
|
||||
func (h *History) List() []Transfer {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
|
||||
list := make([]Transfer, len(h.entries))
|
||||
for i, e := range h.entries {
|
||||
list[i] = e.clone()
|
||||
}
|
||||
return list
|
||||
}
|
||||
|
||||
// Delete removes one entry and persists the log.
|
||||
func (h *History) Delete(id OfferID) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
if i := h.indexOf(id); i >= 0 {
|
||||
h.entries = slices.Delete(h.entries, i, i+1)
|
||||
h.persist()
|
||||
}
|
||||
}
|
||||
|
||||
func (h *History) indexOf(id OfferID) int {
|
||||
return slices.IndexFunc(h.entries, func(t Transfer) bool { return t.ID == id })
|
||||
}
|
||||
|
||||
func (h *History) prune() {
|
||||
if len(h.entries) <= historyCap {
|
||||
return
|
||||
}
|
||||
for i := len(h.entries) - 1; i >= 0 && len(h.entries) > historyCap; i-- {
|
||||
if h.entries[i].terminal() {
|
||||
h.entries = slices.Delete(h.entries, i, i+1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// settleInterrupted closes out transfers that were still moving when the
|
||||
// process died. Nothing is left to finish them, so left alone they would sit in
|
||||
// the log as permanently pending. Reports whether anything changed.
|
||||
func (h *History) settleInterrupted() bool {
|
||||
changed := false
|
||||
for i, t := range h.entries {
|
||||
if t.terminal() {
|
||||
continue
|
||||
}
|
||||
h.entries[i].State = StateFailed
|
||||
h.entries[i].Reason = ReasonInterrupted
|
||||
h.entries[i].UpdatedAt = time.Now()
|
||||
changed = true
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
func (h *History) persist() {
|
||||
if err := saveSection(h.store, namespaceHistory, h.entries); err != nil {
|
||||
log.Warnf("failed to write file drop history: %v", err)
|
||||
}
|
||||
}
|
||||
221
client/internal/filedrop/http.go
Normal file
221
client/internal/filedrop/http.go
Normal file
@@ -0,0 +1,221 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// httpTransport adapts the receiver to plain HTTP/1.1 over the tunnel. It only
|
||||
// parses requests, maps domain errors to status codes, and encodes responses.
|
||||
type httpTransport struct {
|
||||
recv *receiver
|
||||
}
|
||||
|
||||
func (t *httpTransport) routes() http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc(pathOffers, t.handleOffers)
|
||||
mux.HandleFunc(pathOffersSlash, t.handleOffer)
|
||||
return mux
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleOffers(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||
return
|
||||
}
|
||||
|
||||
sender, ok := t.identify(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var req OfferRequest
|
||||
if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, maxOfferBodySize)).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "malformed offer")
|
||||
return
|
||||
}
|
||||
|
||||
offer, err := t.recv.submitOffer(sender, req)
|
||||
if err != nil {
|
||||
writeDomainError(w, err)
|
||||
return
|
||||
}
|
||||
|
||||
status := http.StatusAccepted
|
||||
if offer.Decision == DecisionAccepted {
|
||||
status = http.StatusCreated
|
||||
}
|
||||
writeJSON(w, status, OfferResponse{ID: offer.ID, Decision: offer.Decision})
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleOffer(w http.ResponseWriter, r *http.Request) {
|
||||
sender, ok := t.identify(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
id, index, hasIndex, err := parseOfferPath(r.URL.Path)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusNotFound, "unknown path")
|
||||
return
|
||||
}
|
||||
|
||||
if !hasIndex {
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
t.handleOfferStatus(w, r, sender, id)
|
||||
case http.MethodDelete:
|
||||
t.handleOfferCancel(w, sender, id)
|
||||
default:
|
||||
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodPut:
|
||||
t.handleUpload(w, r, sender, id, index)
|
||||
case http.MethodHead:
|
||||
t.handleUploadProbe(w, sender, id, index)
|
||||
default:
|
||||
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||
}
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleOfferStatus(w http.ResponseWriter, r *http.Request, sender senderIdentity, id OfferID) {
|
||||
offer, err := t.recv.awaitDecision(r.Context(), sender, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrOfferNotFound) {
|
||||
writeDomainError(w, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, OfferResponse{ID: offer.ID, Decision: offer.Decision})
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleOfferCancel(w http.ResponseWriter, sender senderIdentity, id OfferID) {
|
||||
if err := t.recv.withdraw(sender, id); err != nil {
|
||||
writeDomainError(w, err)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleUploadProbe(w http.ResponseWriter, sender senderIdentity, id OfferID, index int) {
|
||||
received, err := t.recv.receivedBytes(sender, id, index)
|
||||
if err != nil {
|
||||
writeDomainError(w, err)
|
||||
return
|
||||
}
|
||||
w.Header().Set(HeaderReceivedBytes, strconv.FormatInt(received, 10))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
|
||||
func (t *httpTransport) handleUpload(w http.ResponseWriter, r *http.Request, sender senderIdentity, id OfferID, index int) {
|
||||
offset, err := parseOffset(r.URL.Query().Get("offset"))
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := t.recv.upload(sender, id, index, offset, r.Body); err != nil {
|
||||
writeDomainError(w, err)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func (t *httpTransport) identify(w http.ResponseWriter, r *http.Request) (senderIdentity, bool) {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddr(host)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusForbidden, "unknown sender")
|
||||
return senderIdentity{}, false
|
||||
}
|
||||
|
||||
sender, ok := t.recv.identify(addr)
|
||||
if !ok {
|
||||
writeError(w, http.StatusForbidden, "unknown sender")
|
||||
return senderIdentity{}, false
|
||||
}
|
||||
return sender, true
|
||||
}
|
||||
|
||||
func writeDomainError(w http.ResponseWriter, err error) {
|
||||
switch {
|
||||
case errors.Is(err, ErrRefused), errors.Is(err, ErrNotAccepted), errors.Is(err, ErrUnknownPeer):
|
||||
writeError(w, http.StatusForbidden, err.Error())
|
||||
case errors.Is(err, ErrOfferNotFound):
|
||||
writeError(w, http.StatusNotFound, err.Error())
|
||||
case errors.Is(err, ErrInvalidOffer):
|
||||
writeError(w, http.StatusBadRequest, err.Error())
|
||||
case errors.Is(err, ErrStorage):
|
||||
writeError(w, http.StatusInsufficientStorage, err.Error())
|
||||
default:
|
||||
writeError(w, http.StatusInternalServerError, err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func parseOfferPath(path string) (OfferID, int, bool, error) {
|
||||
rest := strings.TrimPrefix(path, pathOffersSlash)
|
||||
if rest == "" || rest == path {
|
||||
return "", 0, false, fmt.Errorf("not an offer path")
|
||||
}
|
||||
|
||||
parts := strings.Split(rest, "/")
|
||||
if parts[0] == "" {
|
||||
return "", 0, false, fmt.Errorf("missing offer id")
|
||||
}
|
||||
id := OfferID(parts[0])
|
||||
|
||||
switch len(parts) {
|
||||
case 1:
|
||||
return id, 0, false, nil
|
||||
case 3:
|
||||
if parts[1] != segmentFiles {
|
||||
return "", 0, false, fmt.Errorf("unknown sub-resource %q", parts[1])
|
||||
}
|
||||
index, err := strconv.Atoi(parts[2])
|
||||
if err != nil || index < 0 {
|
||||
return "", 0, false, fmt.Errorf("invalid file index")
|
||||
}
|
||||
return id, index, true, nil
|
||||
default:
|
||||
return "", 0, false, fmt.Errorf("unknown offer path")
|
||||
}
|
||||
}
|
||||
|
||||
func parseOffset(raw string) (int64, error) {
|
||||
if raw == "" {
|
||||
return 0, nil
|
||||
}
|
||||
offset, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil || offset < 0 {
|
||||
return 0, fmt.Errorf("invalid offset")
|
||||
}
|
||||
return offset, nil
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, status int, body any) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
if err := json.NewEncoder(w).Encode(body); err != nil {
|
||||
log.Debugf("write file drop response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func writeError(w http.ResponseWriter, status int, message string) {
|
||||
writeJSON(w, status, map[string]string{"error": message})
|
||||
}
|
||||
660
client/internal/filedrop/manager.go
Normal file
660
client/internal/filedrop/manager.go
Normal file
@@ -0,0 +1,660 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.zx2c4.com/wireguard/tun/netstack"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
)
|
||||
|
||||
// The event kinds. Progress is not an event: live transfers are polled.
|
||||
const (
|
||||
EventOffer EventKind = iota
|
||||
EventCompleted
|
||||
EventFailed
|
||||
EventWithdrawn
|
||||
)
|
||||
|
||||
// portSignalGrace bounds how long a failed attempt waits for one signal message
|
||||
// that may advertise the receiver's actual port before giving up.
|
||||
const portSignalGrace = 3 * time.Second
|
||||
|
||||
// ErrNotConnected indicates the operation needs a running tunnel.
|
||||
var ErrNotConnected = errors.New("not connected")
|
||||
|
||||
// EventKind classifies the events the manager surfaces to the platform layer.
|
||||
type EventKind uint8
|
||||
|
||||
// EventSink receives transfer events. Calls may come from server goroutines.
|
||||
type EventSink func(kind EventKind, transfer Transfer)
|
||||
|
||||
// ManagerConfig configures a per-profile file drop manager. Policy and history
|
||||
// live in Store; DataDir only holds the spool of partially received files,
|
||||
// which is disposable and never outlives an offer's TTL.
|
||||
type ManagerConfig struct {
|
||||
Profile profilemanager.ID
|
||||
DataDir string
|
||||
Store Store
|
||||
Events EventSink
|
||||
OfferTTL time.Duration
|
||||
}
|
||||
|
||||
type sendHandle struct {
|
||||
cancel context.CancelFunc
|
||||
ip netip.Addr
|
||||
addr netip.AddrPort
|
||||
remoteID OfferID
|
||||
}
|
||||
|
||||
// Manager owns one profile's file drop state.
|
||||
type Manager struct {
|
||||
mu sync.Mutex
|
||||
profile profilemanager.ID
|
||||
dataDir string
|
||||
store Store
|
||||
policy *PolicyStore
|
||||
history *History
|
||||
events EventSink
|
||||
offerTTL time.Duration
|
||||
|
||||
server *Server
|
||||
ports *PortRegistry
|
||||
dial DialFunc
|
||||
senderName string
|
||||
sends map[OfferID]*sendHandle
|
||||
sendWg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewManager loads or initializes the file drop state for one profile.
|
||||
func NewManager(cfg ManagerConfig) (*Manager, error) {
|
||||
if cfg.DataDir == "" {
|
||||
return nil, errors.New("data dir is required")
|
||||
}
|
||||
if err := os.MkdirAll(cfg.DataDir, 0o700); err != nil {
|
||||
return nil, fmt.Errorf("create file drop dir: %w", err)
|
||||
}
|
||||
|
||||
m := &Manager{
|
||||
profile: cfg.Profile,
|
||||
dataDir: cfg.DataDir,
|
||||
store: cfg.Store,
|
||||
policy: LoadPolicyStore(cfg.Profile, cfg.Store),
|
||||
history: LoadHistory(cfg.Store),
|
||||
events: cfg.Events,
|
||||
offerTTL: cfg.OfferTTL,
|
||||
ports: NewPortRegistry(),
|
||||
sends: make(map[OfferID]*sendHandle),
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// Profile returns the profile this manager belongs to.
|
||||
func (m *Manager) Profile() profilemanager.ID {
|
||||
return m.profile
|
||||
}
|
||||
|
||||
// Policy returns the receiving policy store.
|
||||
func (m *Manager) Policy() *PolicyStore {
|
||||
return m.policy
|
||||
}
|
||||
|
||||
// Ports returns the registry of peer-advertised listen ports; the engine feeds it
|
||||
// from incoming signal messages.
|
||||
func (m *Manager) Ports() *PortRegistry {
|
||||
return m.ports
|
||||
}
|
||||
|
||||
// ReceiverPort returns the port the receiver is actually bound to, 0 when stopped.
|
||||
func (m *Manager) ReceiverPort() uint16 {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
if server == nil {
|
||||
return 0
|
||||
}
|
||||
return server.BoundPort()
|
||||
}
|
||||
|
||||
// Transfers returns the history entries, newest first, with pending offers included.
|
||||
func (m *Manager) Transfers() []Transfer {
|
||||
return m.history.List()
|
||||
}
|
||||
|
||||
// DeleteTransfer removes a history entry. A live transfer is cancelled first.
|
||||
func (m *Manager) DeleteTransfer(id OfferID) {
|
||||
if t, ok := m.history.Get(id); ok && !t.terminal() {
|
||||
m.Cancel(id)
|
||||
}
|
||||
m.history.Delete(id)
|
||||
}
|
||||
|
||||
// DestinationDir returns the directory received files are delivered to.
|
||||
func (m *Manager) DestinationDir() string {
|
||||
return m.policy.DestinationDir()
|
||||
}
|
||||
|
||||
// SetDestinationDir persists the delivery directory.
|
||||
func (m *Manager) SetDestinationDir(dir string) error {
|
||||
return m.policy.SetDestinationDir(dir)
|
||||
}
|
||||
|
||||
// StartReceiver binds the receiving server on addr.
|
||||
func (m *Manager) StartReceiver(ctx context.Context, addr netip.AddrPort, netstackNet *netstack.Net, resolver PeerResolver) error {
|
||||
m.mu.Lock()
|
||||
if m.server != nil {
|
||||
m.mu.Unlock()
|
||||
return errors.New("receiver is already running")
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
server, err := NewServer(ServerConfig{
|
||||
SpoolDir: filepath.Join(m.dataDir, "spool"),
|
||||
Policy: m.policy,
|
||||
Resolver: resolver,
|
||||
Notifier: m,
|
||||
OfferTTL: m.offerTTL,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("create receiver: %w", err)
|
||||
}
|
||||
if netstackNet != nil {
|
||||
server.SetNetstackNet(netstackNet)
|
||||
}
|
||||
|
||||
if err := server.Start(ctx, addr); err != nil {
|
||||
return fmt.Errorf("start receiver: %w", err)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
m.server = server
|
||||
m.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddReceiverListener serves the receiver on an additional address, such as IPv6.
|
||||
func (m *Manager) AddReceiverListener(ctx context.Context, addr netip.AddrPort) error {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
|
||||
if server == nil {
|
||||
return errors.New("receiver is not running")
|
||||
}
|
||||
return server.AddListener(ctx, addr)
|
||||
}
|
||||
|
||||
// StopReceiver shuts the receiving server down and drops the tunnel dialer.
|
||||
func (m *Manager) StopReceiver() error {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.server = nil
|
||||
m.dial = nil
|
||||
m.mu.Unlock()
|
||||
|
||||
if server == nil {
|
||||
return nil
|
||||
}
|
||||
return server.Stop()
|
||||
}
|
||||
|
||||
// SetTunnel gives the manager the tunnel dialer and the local sender name.
|
||||
func (m *Manager) SetTunnel(dial DialFunc, senderName string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.dial = dial
|
||||
m.senderName = senderName
|
||||
}
|
||||
|
||||
// Close stops the receiver and aborts every outgoing transfer.
|
||||
func (m *Manager) Close() error {
|
||||
err := m.StopReceiver()
|
||||
|
||||
m.mu.Lock()
|
||||
for _, h := range m.sends {
|
||||
h.cancel()
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
m.sendWg.Wait()
|
||||
return err
|
||||
}
|
||||
|
||||
// Send starts an asynchronous transfer and returns its local transfer ID.
|
||||
func (m *Manager) Send(peer PeerKey, peerName string, addr netip.Addr, payloads []Payload) (OfferID, error) {
|
||||
if len(payloads) == 0 {
|
||||
return "", fmt.Errorf("%w: no payloads", ErrInvalidOffer)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
dial, senderName := m.dial, m.senderName
|
||||
m.mu.Unlock()
|
||||
if dial == nil {
|
||||
return "", ErrNotConnected
|
||||
}
|
||||
|
||||
client, err := NewClient(ClientConfig{Dial: dial, SenderName: senderName, OfferTimeout: m.offerTTL})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
id := OfferID(uuid.NewString())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
handle := &sendHandle{cancel: cancel, ip: addr}
|
||||
|
||||
m.mu.Lock()
|
||||
m.sends[id] = handle
|
||||
m.mu.Unlock()
|
||||
|
||||
transfer := Transfer{
|
||||
ID: id,
|
||||
Direction: DirectionSent,
|
||||
PeerKey: peer,
|
||||
PeerName: peerName,
|
||||
Files: payloadMetas(payloads),
|
||||
State: StatePending,
|
||||
TotalSize: payloadTotal(payloads),
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
m.history.Upsert(transfer)
|
||||
|
||||
m.sendWg.Add(1)
|
||||
go func() {
|
||||
defer m.sendWg.Done()
|
||||
defer cancel()
|
||||
m.runSend(ctx, client, handle, transfer, payloads)
|
||||
|
||||
m.mu.Lock()
|
||||
delete(m.sends, id)
|
||||
m.mu.Unlock()
|
||||
}()
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// Cancel aborts a transfer in either direction.
|
||||
func (m *Manager) Cancel(id OfferID) {
|
||||
m.mu.Lock()
|
||||
handle := m.sends[id]
|
||||
server := m.server
|
||||
var remoteAddr netip.AddrPort
|
||||
var remoteID OfferID
|
||||
if handle != nil {
|
||||
remoteAddr, remoteID = handle.addr, handle.remoteID
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
if handle != nil {
|
||||
handle.cancel()
|
||||
if remoteID != "" && remoteAddr.IsValid() {
|
||||
m.withdrawRemote(remoteAddr, remoteID)
|
||||
}
|
||||
m.finishTransfer(id, StateCancelled, "")
|
||||
return
|
||||
}
|
||||
|
||||
if server != nil {
|
||||
if offer, ok := server.Offers().Decide(id, DecisionDeclined); ok {
|
||||
server.Spool().Remove(offer.ID)
|
||||
}
|
||||
}
|
||||
m.finishTransfer(id, StateCancelled, "")
|
||||
}
|
||||
|
||||
// Accept releases a pending incoming offer for upload.
|
||||
func (m *Manager) Accept(id OfferID) error {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
if server == nil {
|
||||
return ErrNotConnected
|
||||
}
|
||||
|
||||
offer, ok := server.Offers().Decide(id, DecisionAccepted)
|
||||
if !ok {
|
||||
return ErrOfferNotFound
|
||||
}
|
||||
|
||||
if offer.State == StateCompleted {
|
||||
m.OnCompleted(offer)
|
||||
return nil
|
||||
}
|
||||
|
||||
m.history.SetProgress(id, 0)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Decline refuses a pending incoming offer.
|
||||
func (m *Manager) Decline(id OfferID) error {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
if server == nil {
|
||||
return ErrNotConnected
|
||||
}
|
||||
|
||||
offer, ok := server.Offers().Decide(id, DecisionDeclined)
|
||||
if !ok {
|
||||
return ErrOfferNotFound
|
||||
}
|
||||
|
||||
server.Spool().Remove(offer.ID)
|
||||
m.finishTransfer(id, StateDeclined, "")
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetSenderRule records a per-sender exception.
|
||||
func (m *Manager) SetSenderRule(peer PeerKey, rule SenderRule) error {
|
||||
if err := m.policy.SetSenderRule(peer, rule); err != nil {
|
||||
return err
|
||||
}
|
||||
if rule != SenderRuleBlock {
|
||||
return nil
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
if server == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, offer := range server.Offers().List() {
|
||||
if offer.Sender != peer || offer.Decision != DecisionPending {
|
||||
continue
|
||||
}
|
||||
if declined, ok := server.Offers().Decide(offer.ID, DecisionDeclined); ok {
|
||||
server.Spool().Remove(declined.ID)
|
||||
m.finishTransfer(declined.ID, StateDeclined, "")
|
||||
m.emit(EventWithdrawn, m.transferOf(declined.ID))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) runSend(ctx context.Context, client *Client, handle *sendHandle, transfer Transfer, payloads []Payload) {
|
||||
addr, remoteID, decision, err := m.offerWithPortRetry(ctx, client, handle, transfer.PeerKey, payloads)
|
||||
if err != nil {
|
||||
m.failSend(ctx, transfer.ID, err)
|
||||
return
|
||||
}
|
||||
|
||||
decision, err = client.AwaitDecision(ctx, addr, remoteID, decision)
|
||||
if err != nil {
|
||||
m.failSend(ctx, transfer.ID, err)
|
||||
return
|
||||
}
|
||||
if err := decisionError(decision); err != nil {
|
||||
m.failSend(ctx, transfer.ID, err)
|
||||
return
|
||||
}
|
||||
|
||||
m.history.SetProgress(transfer.ID, 0)
|
||||
|
||||
completed := make([]int64, len(payloads))
|
||||
progress := func(index int, sent, _ int64) {
|
||||
completed[index] = sent
|
||||
var total int64
|
||||
for _, n := range completed {
|
||||
total += n
|
||||
}
|
||||
m.history.SetProgress(transfer.ID, total)
|
||||
}
|
||||
|
||||
if err := client.Upload(ctx, addr, remoteID, payloads, progress); err != nil {
|
||||
m.failSend(ctx, transfer.ID, err)
|
||||
return
|
||||
}
|
||||
|
||||
m.history.SetProgress(transfer.ID, transfer.TotalSize)
|
||||
m.finishTransfer(transfer.ID, StateCompleted, "")
|
||||
m.emit(EventCompleted, m.transferOf(transfer.ID))
|
||||
}
|
||||
|
||||
// offerWithPortRetry places the offer on the last advertised port, falling back to
|
||||
// the default. When the attempt fails on the transport, it waits out one signal
|
||||
// message that may carry the receiver's actual port and retries there once. A port
|
||||
// learned mid-attempt aborts the attempt immediately instead of letting it hang.
|
||||
func (m *Manager) offerWithPortRetry(ctx context.Context, client *Client, handle *sendHandle, key PeerKey, payloads []Payload) (netip.AddrPort, OfferID, Decision, error) {
|
||||
used := m.ports.Port(key)
|
||||
addr := netip.AddrPortFrom(handle.ip, effectivePort(used))
|
||||
|
||||
remoteID, decision, err := m.offerWatchingPorts(ctx, client, key, used, addr, payloads)
|
||||
if err == nil {
|
||||
m.storeRemote(handle, addr, remoteID)
|
||||
return addr, remoteID, decision, nil
|
||||
}
|
||||
if ctx.Err() != nil || !transportFailure(err) {
|
||||
return addr, remoteID, decision, err
|
||||
}
|
||||
|
||||
graceCtx, cancel := context.WithTimeout(ctx, portSignalGrace)
|
||||
port, changed := m.ports.Await(graceCtx, key, used)
|
||||
cancel()
|
||||
if !changed {
|
||||
return addr, remoteID, decision, err
|
||||
}
|
||||
|
||||
addr = netip.AddrPortFrom(handle.ip, effectivePort(port))
|
||||
remoteID, decision, err = client.Offer(ctx, addr, payloads)
|
||||
if err != nil {
|
||||
return addr, remoteID, decision, err
|
||||
}
|
||||
m.storeRemote(handle, addr, remoteID)
|
||||
return addr, remoteID, decision, nil
|
||||
}
|
||||
|
||||
// offerWatchingPorts runs the offer while watching for a port advertisement that
|
||||
// differs from the one in use; such an advertisement aborts the in-flight attempt.
|
||||
func (m *Manager) offerWatchingPorts(ctx context.Context, client *Client, key PeerKey, used uint16, addr netip.AddrPort, payloads []Payload) (OfferID, Decision, error) {
|
||||
watchCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
go func() {
|
||||
if _, changed := m.ports.Await(watchCtx, key, used); changed {
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
|
||||
return client.Offer(watchCtx, addr, payloads)
|
||||
}
|
||||
|
||||
func (m *Manager) storeRemote(handle *sendHandle, addr netip.AddrPort, remoteID OfferID) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
handle.addr = addr
|
||||
handle.remoteID = remoteID
|
||||
}
|
||||
|
||||
func (m *Manager) failSend(ctx context.Context, id OfferID, err error) {
|
||||
if ctx.Err() != nil {
|
||||
m.finishTransfer(id, StateCancelled, "")
|
||||
return
|
||||
}
|
||||
|
||||
state := StateFailed
|
||||
switch {
|
||||
case errors.Is(err, ErrDeclined):
|
||||
state = StateDeclined
|
||||
case errors.Is(err, ErrExpired):
|
||||
state = StateExpired
|
||||
}
|
||||
|
||||
message := ""
|
||||
reason := ReasonNone
|
||||
if state == StateFailed {
|
||||
message = err.Error()
|
||||
if transportFailure(err) {
|
||||
reason = ReasonUnreachable
|
||||
}
|
||||
}
|
||||
m.finishTransferReason(id, state, message, reason)
|
||||
m.emit(EventFailed, m.transferOf(id))
|
||||
}
|
||||
|
||||
func (m *Manager) withdrawRemote(addr netip.AddrPort, remoteID OfferID) {
|
||||
m.mu.Lock()
|
||||
dial, senderName := m.dial, m.senderName
|
||||
m.mu.Unlock()
|
||||
if dial == nil {
|
||||
return
|
||||
}
|
||||
|
||||
client, err := NewClient(ClientConfig{Dial: dial, SenderName: senderName})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
if err := client.Cancel(ctx, addr, remoteID); err != nil {
|
||||
log.Debugf("failed to withdraw file drop offer: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// OnOffer implements Notifier for the receiver server.
|
||||
func (m *Manager) OnOffer(offer Offer) {
|
||||
transfer := Transfer{
|
||||
ID: offer.ID,
|
||||
Direction: DirectionReceived,
|
||||
PeerKey: offer.Sender,
|
||||
PeerName: offer.SenderName,
|
||||
Files: offer.Files,
|
||||
State: offer.State,
|
||||
TotalSize: offer.TotalSize(),
|
||||
CreatedAt: offer.CreatedAt,
|
||||
}
|
||||
m.history.Upsert(transfer)
|
||||
|
||||
if offer.Decision == DecisionPending {
|
||||
m.emit(EventOffer, transfer)
|
||||
}
|
||||
if offer.Decision == DecisionAccepted && offer.State == StateCompleted {
|
||||
m.OnCompleted(offer)
|
||||
}
|
||||
}
|
||||
|
||||
// OnProgress implements Notifier.
|
||||
func (m *Manager) OnProgress(offer Offer, index int, received int64) {
|
||||
var total int64
|
||||
for i, n := range offer.Progress {
|
||||
if i == index {
|
||||
n = received
|
||||
}
|
||||
total += n
|
||||
}
|
||||
m.history.SetProgress(offer.ID, total)
|
||||
}
|
||||
|
||||
// OnCompleted implements Notifier.
|
||||
func (m *Manager) OnCompleted(offer Offer) {
|
||||
m.mu.Lock()
|
||||
server := m.server
|
||||
m.mu.Unlock()
|
||||
if server == nil {
|
||||
return
|
||||
}
|
||||
|
||||
transfer, ok := m.history.Get(offer.ID)
|
||||
if !ok || transfer.State == StateCompleted {
|
||||
return
|
||||
}
|
||||
|
||||
delivered, err := deliver(server.Spool(), offer, m.policy.DestinationDir())
|
||||
if err != nil {
|
||||
log.Errorf("failed to deliver file drop payloads: %v", err)
|
||||
m.finishTransfer(offer.ID, StateFailed, err.Error())
|
||||
m.emit(EventFailed, m.transferOf(offer.ID))
|
||||
return
|
||||
}
|
||||
|
||||
transfer.State = StateCompleted
|
||||
transfer.Transferred = transfer.TotalSize
|
||||
transfer.DeliveredPaths = delivered
|
||||
transfer.Error = ""
|
||||
m.history.Upsert(transfer)
|
||||
m.emit(EventCompleted, transfer)
|
||||
}
|
||||
|
||||
// OnFailed implements Notifier.
|
||||
func (m *Manager) OnFailed(offer Offer, err error) {
|
||||
if errors.Is(err, ErrExpired) {
|
||||
m.finishTransfer(offer.ID, StateExpired, "")
|
||||
m.emit(EventWithdrawn, m.transferOf(offer.ID))
|
||||
return
|
||||
}
|
||||
m.finishTransfer(offer.ID, StateFailed, err.Error())
|
||||
m.emit(EventFailed, m.transferOf(offer.ID))
|
||||
}
|
||||
|
||||
// OnWithdrawn implements Notifier: the sender cancelled, so the consent prompt goes away.
|
||||
func (m *Manager) OnWithdrawn(offer Offer) {
|
||||
m.finishTransfer(offer.ID, StateCancelled, "")
|
||||
m.emit(EventWithdrawn, m.transferOf(offer.ID))
|
||||
}
|
||||
|
||||
func (m *Manager) finishTransfer(id OfferID, state State, message string) {
|
||||
m.finishTransferReason(id, state, message, ReasonNone)
|
||||
}
|
||||
|
||||
func (m *Manager) finishTransferReason(id OfferID, state State, message string, reason FailureReason) {
|
||||
transfer, ok := m.history.Get(id)
|
||||
if !ok || transfer.terminal() {
|
||||
return
|
||||
}
|
||||
transfer.State = state
|
||||
transfer.Error = message
|
||||
transfer.Reason = reason
|
||||
m.history.Upsert(transfer)
|
||||
}
|
||||
|
||||
func (m *Manager) transferOf(id OfferID) Transfer {
|
||||
t, _ := m.history.Get(id)
|
||||
return t
|
||||
}
|
||||
|
||||
func (m *Manager) emit(kind EventKind, transfer Transfer) {
|
||||
if m.events != nil && transfer.ID != "" {
|
||||
m.events(kind, transfer)
|
||||
}
|
||||
}
|
||||
|
||||
func payloadMetas(payloads []Payload) []FileMeta {
|
||||
metas := make([]FileMeta, len(payloads))
|
||||
for i, p := range payloads {
|
||||
metas[i] = p.Meta
|
||||
}
|
||||
return metas
|
||||
}
|
||||
|
||||
func payloadTotal(payloads []Payload) int64 {
|
||||
var total int64
|
||||
for _, p := range payloads {
|
||||
total += p.Meta.Size
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func effectivePort(advertised uint16) uint16 {
|
||||
if advertised == 0 {
|
||||
return Port
|
||||
}
|
||||
return advertised
|
||||
}
|
||||
|
||||
// transportFailure reports whether the offer never reached the receiver; any HTTP
|
||||
// response, refusal included, proves the port right and is not retried elsewhere.
|
||||
func transportFailure(err error) bool {
|
||||
var urlErr *url.Error
|
||||
return errors.As(err, &urlErr)
|
||||
}
|
||||
282
client/internal/filedrop/offer.go
Normal file
282
client/internal/filedrop/offer.go
Normal file
@@ -0,0 +1,282 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Offer is one incoming transfer as tracked by the receiver.
|
||||
type Offer struct {
|
||||
ID OfferID
|
||||
Sender PeerKey
|
||||
SenderName string
|
||||
Files []FileMeta
|
||||
Decision Decision
|
||||
State State
|
||||
CreatedAt time.Time
|
||||
ExpiresAt time.Time
|
||||
Progress []int64
|
||||
}
|
||||
|
||||
type offerEntry struct {
|
||||
offer Offer
|
||||
decided chan struct{}
|
||||
}
|
||||
|
||||
// OfferStore tracks incoming offers and their decisions.
|
||||
type OfferStore struct {
|
||||
mu sync.RWMutex
|
||||
offers map[OfferID]*offerEntry
|
||||
ttl time.Duration
|
||||
newID func() OfferID
|
||||
nowFunc func() time.Time
|
||||
}
|
||||
|
||||
// NewOfferStore returns an empty store using ttl as the decision deadline.
|
||||
func NewOfferStore(ttl time.Duration) *OfferStore {
|
||||
if ttl <= 0 {
|
||||
ttl = DefaultOfferTTL
|
||||
}
|
||||
return &OfferStore{
|
||||
offers: make(map[OfferID]*offerEntry),
|
||||
ttl: ttl,
|
||||
newID: func() OfferID { return OfferID(uuid.NewString()) },
|
||||
nowFunc: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
// Add registers a new offer with the given initial decision and returns its snapshot.
|
||||
func (s *OfferStore) Add(sender PeerKey, senderName string, files []FileMeta, decision Decision) Offer {
|
||||
now := s.nowFunc()
|
||||
|
||||
entry := &offerEntry{
|
||||
offer: Offer{
|
||||
ID: s.newID(),
|
||||
Sender: sender,
|
||||
SenderName: senderName,
|
||||
Files: files,
|
||||
Decision: decision,
|
||||
State: stateForDecision(decision),
|
||||
CreatedAt: now,
|
||||
ExpiresAt: now.Add(s.ttl),
|
||||
Progress: make([]int64, len(files)),
|
||||
},
|
||||
decided: make(chan struct{}),
|
||||
}
|
||||
if decision != DecisionPending {
|
||||
close(entry.decided)
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.offers[entry.offer.ID] = entry
|
||||
s.mu.Unlock()
|
||||
|
||||
return entry.offer.clone()
|
||||
}
|
||||
|
||||
// Get returns a snapshot of one offer belonging to sender.
|
||||
func (s *OfferStore) Get(sender PeerKey, id OfferID) (Offer, bool) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
entry, ok := s.offers[id]
|
||||
if !ok || entry.offer.Sender != sender {
|
||||
return Offer{}, false
|
||||
}
|
||||
return entry.offer.clone(), true
|
||||
}
|
||||
|
||||
// List returns snapshots of every tracked offer.
|
||||
func (s *OfferStore) List() []Offer {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
offers := make([]Offer, 0, len(s.offers))
|
||||
for _, entry := range s.offers {
|
||||
offers = append(offers, entry.offer.clone())
|
||||
}
|
||||
return offers
|
||||
}
|
||||
|
||||
// Decide records the receiver's answer; a made decision is final.
|
||||
func (s *OfferStore) Decide(id OfferID, decision Decision) (Offer, bool) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
entry, ok := s.offers[id]
|
||||
if !ok || entry.offer.Decision != DecisionPending {
|
||||
return Offer{}, false
|
||||
}
|
||||
|
||||
entry.offer.Decision = decision
|
||||
entry.offer.State = stateForDecision(decision)
|
||||
if decision == DecisionAccepted && entry.offer.awaitsNoUpload() {
|
||||
entry.offer.State = StateCompleted
|
||||
}
|
||||
close(entry.decided)
|
||||
|
||||
return entry.offer.clone(), true
|
||||
}
|
||||
|
||||
// Await blocks until a decision, expiry, or ctx cancellation.
|
||||
func (s *OfferStore) Await(ctx context.Context, sender PeerKey, id OfferID) (Offer, error) {
|
||||
s.mu.RLock()
|
||||
entry, ok := s.offers[id]
|
||||
if ok && entry.offer.Sender != sender {
|
||||
ok = false
|
||||
}
|
||||
var decided chan struct{}
|
||||
var expiresAt time.Time
|
||||
if ok {
|
||||
decided = entry.decided
|
||||
expiresAt = entry.offer.ExpiresAt
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
if !ok {
|
||||
return Offer{}, ErrOfferNotFound
|
||||
}
|
||||
|
||||
timer := time.NewTimer(time.Until(expiresAt))
|
||||
defer timer.Stop()
|
||||
|
||||
select {
|
||||
case <-decided:
|
||||
case <-timer.C:
|
||||
s.Decide(id, DecisionExpired)
|
||||
case <-ctx.Done():
|
||||
offer, _ := s.Get(sender, id)
|
||||
return offer, ctx.Err()
|
||||
}
|
||||
|
||||
offer, ok := s.Get(sender, id)
|
||||
if !ok {
|
||||
return Offer{}, ErrOfferNotFound
|
||||
}
|
||||
return offer, nil
|
||||
}
|
||||
|
||||
// SetProgress records the staged byte count for one file of an offer.
|
||||
func (s *OfferStore) SetProgress(id OfferID, index int, received int64) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
entry, ok := s.offers[id]
|
||||
if !ok || index < 0 || index >= len(entry.offer.Progress) {
|
||||
return
|
||||
}
|
||||
entry.offer.Progress[index] = received
|
||||
if entry.offer.State == StatePending {
|
||||
entry.offer.State = StateTransferring
|
||||
}
|
||||
}
|
||||
|
||||
// SetState overrides the transfer state, for completion and failure reporting.
|
||||
func (s *OfferStore) SetState(id OfferID, state State) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if entry, ok := s.offers[id]; ok {
|
||||
entry.offer.State = state
|
||||
}
|
||||
}
|
||||
|
||||
// Remove drops an offer from the store.
|
||||
func (s *OfferStore) Remove(id OfferID) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.offers, id)
|
||||
}
|
||||
|
||||
// ExpireOverdue marks every pending offer past its deadline as expired and returns them.
|
||||
func (s *OfferStore) ExpireOverdue() []Offer {
|
||||
now := s.nowFunc()
|
||||
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
var expired []Offer
|
||||
for _, entry := range s.offers {
|
||||
if entry.offer.Decision != DecisionPending || now.Before(entry.offer.ExpiresAt) {
|
||||
continue
|
||||
}
|
||||
entry.offer.Decision = DecisionExpired
|
||||
entry.offer.State = StateExpired
|
||||
close(entry.decided)
|
||||
expired = append(expired, entry.offer.clone())
|
||||
}
|
||||
return expired
|
||||
}
|
||||
|
||||
// Complete marks an offer completed once every file reached its announced size.
|
||||
func (s *OfferStore) Complete(id OfferID) (Offer, bool) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
entry, ok := s.offers[id]
|
||||
if !ok || entry.offer.State == StateCompleted {
|
||||
return Offer{}, false
|
||||
}
|
||||
|
||||
if !entry.offer.fullyStaged() {
|
||||
return Offer{}, false
|
||||
}
|
||||
|
||||
entry.offer.State = StateCompleted
|
||||
return entry.offer.clone(), true
|
||||
}
|
||||
|
||||
func (o Offer) awaitsNoUpload() bool {
|
||||
for _, f := range o.Files {
|
||||
if f.Kind != KindText {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (o Offer) fullyStaged() bool {
|
||||
for i, f := range o.Files {
|
||||
if f.Kind == KindText {
|
||||
continue
|
||||
}
|
||||
if i >= len(o.Progress) || o.Progress[i] < f.Size {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (o Offer) clone() Offer {
|
||||
c := o
|
||||
c.Files = make([]FileMeta, len(o.Files))
|
||||
copy(c.Files, o.Files)
|
||||
c.Progress = make([]int64, len(o.Progress))
|
||||
copy(c.Progress, o.Progress)
|
||||
return c
|
||||
}
|
||||
|
||||
// TotalSize is the announced byte count across every file of the offer.
|
||||
func (o Offer) TotalSize() int64 {
|
||||
var total int64
|
||||
for _, f := range o.Files {
|
||||
total += f.Size
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func stateForDecision(d Decision) State {
|
||||
switch d {
|
||||
case DecisionAccepted:
|
||||
return StateTransferring
|
||||
case DecisionDeclined:
|
||||
return StateDeclined
|
||||
case DecisionExpired:
|
||||
return StateExpired
|
||||
default:
|
||||
return StatePending
|
||||
}
|
||||
}
|
||||
206
client/internal/filedrop/policy.go
Normal file
206
client/internal/filedrop/policy.go
Normal file
@@ -0,0 +1,206 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
)
|
||||
|
||||
const (
|
||||
SenderRuleDefault SenderRule = iota
|
||||
SenderRuleAlwaysAccept
|
||||
SenderRuleBlock
|
||||
)
|
||||
|
||||
// SenderRule is a per-sender override on top of the base mode.
|
||||
type SenderRule uint8
|
||||
|
||||
// Policy is the device-local receiving policy of one profile. An empty
|
||||
// DestinationDir means the platform's default download directory.
|
||||
type Policy struct {
|
||||
Mode Mode `json:"mode"`
|
||||
Senders map[PeerKey]SenderRule `json:"senders,omitempty"`
|
||||
DestinationDir string `json:"destinationDir,omitempty"`
|
||||
}
|
||||
|
||||
// PolicyStore holds the receiving policy of one profile and evaluates it per sender.
|
||||
type PolicyStore struct {
|
||||
mu sync.RWMutex
|
||||
profile profilemanager.ID
|
||||
policy Policy
|
||||
store Store
|
||||
}
|
||||
|
||||
// NewPolicyStore returns an in-memory store seeded with the default policy.
|
||||
func NewPolicyStore(profile profilemanager.ID) *PolicyStore {
|
||||
return &PolicyStore{profile: profile, policy: DefaultPolicy()}
|
||||
}
|
||||
|
||||
// LoadPolicyStore builds a store from the persisted policy of one profile.
|
||||
func LoadPolicyStore(profile profilemanager.ID, store Store) *PolicyStore {
|
||||
s := &PolicyStore{
|
||||
profile: profile,
|
||||
policy: DefaultPolicy(),
|
||||
store: store,
|
||||
}
|
||||
if store == nil {
|
||||
return s
|
||||
}
|
||||
|
||||
policy := DefaultPolicy()
|
||||
if err := loadSection(store, namespacePolicy, &policy); err != nil {
|
||||
log.Warnf("failed to load file drop policy for profile %s, using defaults: %v", profile, err)
|
||||
return s
|
||||
}
|
||||
if err := policy.validate(); err != nil {
|
||||
log.Warnf("stored file drop policy for profile %s is invalid, using defaults: %v", profile, err)
|
||||
return s
|
||||
}
|
||||
|
||||
s.policy = policy.normalized()
|
||||
return s
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (r SenderRule) String() string {
|
||||
switch r {
|
||||
case SenderRuleDefault:
|
||||
return "default"
|
||||
case SenderRuleAlwaysAccept:
|
||||
return "always"
|
||||
case SenderRuleBlock:
|
||||
return "block"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(r))
|
||||
}
|
||||
}
|
||||
|
||||
func (p Policy) validate() error {
|
||||
if !p.Mode.valid() {
|
||||
return fmt.Errorf("invalid mode %s", p.Mode)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p Policy) normalized() Policy {
|
||||
c := p.clone()
|
||||
if c.Senders == nil {
|
||||
c.Senders = map[PeerKey]SenderRule{}
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (p Policy) clone() Policy {
|
||||
c := p
|
||||
c.Senders = make(map[PeerKey]SenderRule, len(p.Senders))
|
||||
for k, v := range p.Senders {
|
||||
c.Senders[k] = v
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// Profile returns the profile this policy belongs to.
|
||||
func (s *PolicyStore) Profile() profilemanager.ID {
|
||||
return s.profile
|
||||
}
|
||||
|
||||
// Get returns a copy of the current policy.
|
||||
func (s *PolicyStore) Get() Policy {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.policy.clone()
|
||||
}
|
||||
|
||||
// Set replaces the policy and persists it.
|
||||
func (s *PolicyStore) Set(p Policy) error {
|
||||
if err := p.validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.policy = p.normalized()
|
||||
store, stored := s.store, s.policy.clone()
|
||||
s.mu.Unlock()
|
||||
|
||||
return saveSection(store, namespacePolicy, stored)
|
||||
}
|
||||
|
||||
// SetMode changes the base mode, leaving per-sender rules untouched.
|
||||
func (s *PolicyStore) SetMode(m Mode) error {
|
||||
if !m.valid() {
|
||||
return fmt.Errorf("invalid mode %s", m)
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.policy.Mode = m
|
||||
store, stored := s.store, s.policy.clone()
|
||||
s.mu.Unlock()
|
||||
|
||||
return saveSection(store, namespacePolicy, stored)
|
||||
}
|
||||
|
||||
// SetSenderRule sets or clears the override for a single sender.
|
||||
func (s *PolicyStore) SetSenderRule(key PeerKey, rule SenderRule) error {
|
||||
s.mu.Lock()
|
||||
if rule == SenderRuleDefault {
|
||||
delete(s.policy.Senders, key)
|
||||
} else {
|
||||
if s.policy.Senders == nil {
|
||||
s.policy.Senders = map[PeerKey]SenderRule{}
|
||||
}
|
||||
s.policy.Senders[key] = rule
|
||||
}
|
||||
store, stored := s.store, s.policy.clone()
|
||||
s.mu.Unlock()
|
||||
|
||||
return saveSection(store, namespacePolicy, stored)
|
||||
}
|
||||
|
||||
// DestinationDir returns the directory received files are delivered to.
|
||||
func (s *PolicyStore) DestinationDir() string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.policy.DestinationDir
|
||||
}
|
||||
|
||||
// SetDestinationDir persists the delivery directory.
|
||||
func (s *PolicyStore) SetDestinationDir(dir string) error {
|
||||
s.mu.Lock()
|
||||
s.policy.DestinationDir = dir
|
||||
store, stored := s.store, s.policy.clone()
|
||||
s.mu.Unlock()
|
||||
|
||||
return saveSection(store, namespacePolicy, stored)
|
||||
}
|
||||
|
||||
// Evaluate returns the mode that applies to one sender, denying on unknown values.
|
||||
func (s *PolicyStore) Evaluate(key PeerKey) Mode {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
switch s.policy.Senders[key] {
|
||||
case SenderRuleBlock:
|
||||
return ModeOff
|
||||
case SenderRuleAlwaysAccept:
|
||||
return ModeAutoAccept
|
||||
case SenderRuleDefault:
|
||||
default:
|
||||
return ModeOff
|
||||
}
|
||||
|
||||
if !s.policy.Mode.valid() {
|
||||
return ModeOff
|
||||
}
|
||||
return s.policy.Mode
|
||||
}
|
||||
|
||||
// DefaultPolicy asks before accepting anything, so receiving is never silently on.
|
||||
func DefaultPolicy() Policy {
|
||||
return Policy{
|
||||
Mode: ModeAsk,
|
||||
Senders: map[PeerKey]SenderRule{},
|
||||
}
|
||||
}
|
||||
81
client/internal/filedrop/ports.go
Normal file
81
client/internal/filedrop/ports.go
Normal file
@@ -0,0 +1,81 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// PortRegistry tracks the file drop listen port each peer advertised over
|
||||
// signaling; 0 means the well-known default. Senders can wait on it to learn a
|
||||
// better port after a failed attempt.
|
||||
type PortRegistry struct {
|
||||
mu sync.Mutex
|
||||
ports map[PeerKey]uint16
|
||||
waits map[PeerKey][]chan uint16
|
||||
}
|
||||
|
||||
// NewPortRegistry returns an empty registry.
|
||||
func NewPortRegistry() *PortRegistry {
|
||||
return &PortRegistry{
|
||||
ports: make(map[PeerKey]uint16),
|
||||
waits: make(map[PeerKey][]chan uint16),
|
||||
}
|
||||
}
|
||||
|
||||
// Set records the port a peer advertised and releases every waiter for it.
|
||||
func (r *PortRegistry) Set(key PeerKey, port uint16) {
|
||||
r.mu.Lock()
|
||||
r.ports[key] = port
|
||||
waiters := r.waits[key]
|
||||
delete(r.waits, key)
|
||||
r.mu.Unlock()
|
||||
|
||||
for _, ch := range waiters {
|
||||
ch <- port
|
||||
}
|
||||
}
|
||||
|
||||
// Port returns the last advertised port for a peer; 0 means default or unknown.
|
||||
func (r *PortRegistry) Port(key PeerKey) uint16 {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return r.ports[key]
|
||||
}
|
||||
|
||||
// Await returns the peer's port as soon as it differs from used, or after the next
|
||||
// advertisement even when it does not, reporting whether it differs. It returns
|
||||
// immediately when the currently known port already differs.
|
||||
func (r *PortRegistry) Await(ctx context.Context, key PeerKey, used uint16) (uint16, bool) {
|
||||
r.mu.Lock()
|
||||
if port, ok := r.ports[key]; ok && port != used {
|
||||
r.mu.Unlock()
|
||||
return port, true
|
||||
}
|
||||
ch := make(chan uint16, 1)
|
||||
r.waits[key] = append(r.waits[key], ch)
|
||||
r.mu.Unlock()
|
||||
|
||||
select {
|
||||
case port := <-ch:
|
||||
return port, port != used
|
||||
case <-ctx.Done():
|
||||
r.drop(key, ch)
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func (r *PortRegistry) drop(key PeerKey, ch chan uint16) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
waiters := r.waits[key]
|
||||
for i, w := range waiters {
|
||||
if w == ch {
|
||||
r.waits[key] = append(waiters[:i], waiters[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
if len(r.waits[key]) == 0 {
|
||||
delete(r.waits, key)
|
||||
}
|
||||
}
|
||||
197
client/internal/filedrop/protocol.go
Normal file
197
client/internal/filedrop/protocol.go
Normal file
@@ -0,0 +1,197 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const Port uint16 = 41421
|
||||
|
||||
// HeaderReceivedBytes carries the receiver's confirmed byte count in a HEAD response.
|
||||
const HeaderReceivedBytes = "Netbird-Received-Bytes"
|
||||
|
||||
// DefaultOfferTTL bounds how long an offer waits for the receiver's decision.
|
||||
const DefaultOfferTTL = 5 * time.Minute
|
||||
|
||||
// MaxOfferFiles bounds the number of items a single offer may announce.
|
||||
const MaxOfferFiles = 512
|
||||
|
||||
// MaxInlineTextSize bounds an inline text snippet, which is held in memory.
|
||||
const MaxInlineTextSize = 64 * 1024
|
||||
|
||||
const maxOfferBodySize = 1 << 20
|
||||
|
||||
const (
|
||||
pathOffers = "/v1/offers"
|
||||
pathOffersSlash = pathOffers + "/"
|
||||
|
||||
segmentFiles = "files"
|
||||
)
|
||||
|
||||
// The decisions an offer can carry. Pending is the only non-final one.
|
||||
const (
|
||||
DecisionPending Decision = iota
|
||||
DecisionAccepted
|
||||
DecisionDeclined
|
||||
DecisionExpired
|
||||
)
|
||||
|
||||
// The receiving modes a profile can be in.
|
||||
const (
|
||||
ModeOff Mode = iota
|
||||
ModeAsk
|
||||
ModeAutoAccept
|
||||
)
|
||||
|
||||
// The payload kinds an offer can announce.
|
||||
const (
|
||||
KindFile Kind = iota
|
||||
KindText
|
||||
)
|
||||
|
||||
// The states a transfer moves through.
|
||||
const (
|
||||
StatePending State = iota
|
||||
StateTransferring
|
||||
StateCompleted
|
||||
StateDeclined
|
||||
StateExpired
|
||||
StateCancelled
|
||||
StateFailed
|
||||
)
|
||||
|
||||
var (
|
||||
ErrOfferNotFound = errors.New("offer not found")
|
||||
ErrRefused = errors.New("offer refused by receiver")
|
||||
ErrDeclined = errors.New("offer declined")
|
||||
ErrExpired = errors.New("offer expired")
|
||||
ErrNotAccepted = errors.New("offer not accepted")
|
||||
ErrUnknownPeer = errors.New("unknown peer")
|
||||
ErrInvalidOffer = errors.New("invalid offer")
|
||||
)
|
||||
|
||||
// OfferID identifies a single transfer offer on the receiving peer.
|
||||
type OfferID string
|
||||
|
||||
// PeerKey is the remote peer's public key, used as the identity for per-sender policy.
|
||||
type PeerKey string
|
||||
|
||||
// Decision is the receiver's answer to an offer.
|
||||
type Decision uint8
|
||||
|
||||
// Mode is the receiver's profile-local policy for incoming offers.
|
||||
type Mode uint8
|
||||
|
||||
// Kind distinguishes payloads that are written to the spool from inline text snippets.
|
||||
type Kind uint8
|
||||
|
||||
// State is the lifecycle state of a transfer, on either side.
|
||||
type State uint8
|
||||
|
||||
// FileMeta describes one payload item announced in an offer.
|
||||
type FileMeta struct {
|
||||
Name string `json:"name"`
|
||||
Size int64 `json:"size"`
|
||||
ContentType string `json:"contentType,omitempty"`
|
||||
Kind Kind `json:"kind,omitempty"`
|
||||
Text string `json:"text,omitempty"`
|
||||
}
|
||||
|
||||
// OfferRequest is the JSON body of POST /v1/offers. It carries metadata only.
|
||||
type OfferRequest struct {
|
||||
SenderName string `json:"senderName,omitempty"`
|
||||
Files []FileMeta `json:"files"`
|
||||
}
|
||||
|
||||
// OfferResponse is returned for an offer and for every poll of its status.
|
||||
type OfferResponse struct {
|
||||
ID OfferID `json:"id"`
|
||||
Decision Decision `json:"decision"`
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (d Decision) String() string {
|
||||
switch d {
|
||||
case DecisionPending:
|
||||
return "pending"
|
||||
case DecisionAccepted:
|
||||
return "accepted"
|
||||
case DecisionDeclined:
|
||||
return "declined"
|
||||
case DecisionExpired:
|
||||
return "expired"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(d))
|
||||
}
|
||||
}
|
||||
|
||||
func (d Decision) valid() bool {
|
||||
switch d {
|
||||
case DecisionPending, DecisionAccepted, DecisionDeclined, DecisionExpired:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (m Mode) String() string {
|
||||
switch m {
|
||||
case ModeOff:
|
||||
return "off"
|
||||
case ModeAsk:
|
||||
return "ask"
|
||||
case ModeAutoAccept:
|
||||
return "auto"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(m))
|
||||
}
|
||||
}
|
||||
|
||||
func (m Mode) valid() bool {
|
||||
switch m {
|
||||
case ModeOff, ModeAsk, ModeAutoAccept:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (k Kind) String() string {
|
||||
switch k {
|
||||
case KindFile:
|
||||
return "file"
|
||||
case KindText:
|
||||
return "text"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(k))
|
||||
}
|
||||
}
|
||||
|
||||
func (k Kind) valid() bool {
|
||||
return k == KindFile || k == KindText
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (s State) String() string {
|
||||
switch s {
|
||||
case StatePending:
|
||||
return "pending"
|
||||
case StateTransferring:
|
||||
return "transferring"
|
||||
case StateCompleted:
|
||||
return "completed"
|
||||
case StateDeclined:
|
||||
return "declined"
|
||||
case StateExpired:
|
||||
return "expired"
|
||||
case StateCancelled:
|
||||
return "cancelled"
|
||||
case StateFailed:
|
||||
return "failed"
|
||||
default:
|
||||
return fmt.Sprintf("unknown(%d)", uint8(s))
|
||||
}
|
||||
}
|
||||
226
client/internal/filedrop/receiver.go
Normal file
226
client/internal/filedrop/receiver.go
Normal file
@@ -0,0 +1,226 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// ErrStorage indicates the receiver could not stage payload data locally.
|
||||
var ErrStorage = errors.New("storage failure")
|
||||
|
||||
type senderIdentity struct {
|
||||
key PeerKey
|
||||
name string
|
||||
}
|
||||
|
||||
// receiver implements the transfer protocol independent of any transport. Every
|
||||
// operation takes the already-authenticated sender identity and returns domain
|
||||
// errors for the transport to map.
|
||||
type receiver struct {
|
||||
policy *PolicyStore
|
||||
resolver PeerResolver
|
||||
notifier Notifier
|
||||
offers *OfferStore
|
||||
spool *Spool
|
||||
spoolMaxAge time.Duration
|
||||
}
|
||||
|
||||
func newReceiver(cfg ServerConfig, spool *Spool, maxAge time.Duration) *receiver {
|
||||
return &receiver{
|
||||
policy: cfg.Policy,
|
||||
resolver: cfg.Resolver,
|
||||
notifier: cfg.Notifier,
|
||||
offers: NewOfferStore(cfg.OfferTTL),
|
||||
spool: spool,
|
||||
spoolMaxAge: maxAge,
|
||||
}
|
||||
}
|
||||
|
||||
// identify maps a source overlay address to a known peer, refusing unknown ones.
|
||||
func (r *receiver) identify(addr netip.Addr) (senderIdentity, bool) {
|
||||
key, name, ok := r.resolver.ResolvePeer(addr.Unmap())
|
||||
if !ok {
|
||||
return senderIdentity{}, false
|
||||
}
|
||||
return senderIdentity{key: key, name: name}, true
|
||||
}
|
||||
|
||||
func (r *receiver) submitOffer(sender senderIdentity, req OfferRequest) (Offer, error) {
|
||||
if err := validateOffer(req.Files); err != nil {
|
||||
return Offer{}, fmt.Errorf("%w: %s", ErrInvalidOffer, err)
|
||||
}
|
||||
|
||||
mode := r.policy.Evaluate(sender.key)
|
||||
if mode == ModeOff {
|
||||
return Offer{}, ErrRefused
|
||||
}
|
||||
|
||||
senderName := sender.name
|
||||
if senderName == "" {
|
||||
senderName = req.SenderName
|
||||
}
|
||||
|
||||
decision := DecisionPending
|
||||
if mode == ModeAutoAccept {
|
||||
decision = DecisionAccepted
|
||||
}
|
||||
|
||||
offer := r.offers.Add(sender.key, senderName, req.Files, decision)
|
||||
if err := r.spool.Prepare(offer.ID); err != nil {
|
||||
r.offers.Remove(offer.ID)
|
||||
log.Errorf("prepare spool for offer: %v", err)
|
||||
return Offer{}, fmt.Errorf("%w: prepare spool", ErrStorage)
|
||||
}
|
||||
|
||||
r.notifyOffer(offer)
|
||||
|
||||
if offer.Decision == DecisionAccepted {
|
||||
if completed, done := r.offers.Complete(offer.ID); done {
|
||||
r.notifyCompleted(completed)
|
||||
}
|
||||
}
|
||||
return offer, nil
|
||||
}
|
||||
|
||||
func (r *receiver) awaitDecision(ctx context.Context, sender senderIdentity, id OfferID) (Offer, error) {
|
||||
return r.offers.Await(ctx, sender.key, id)
|
||||
}
|
||||
|
||||
func (r *receiver) withdraw(sender senderIdentity, id OfferID) error {
|
||||
offer, ok := r.offers.Get(sender.key, id)
|
||||
if !ok {
|
||||
return ErrOfferNotFound
|
||||
}
|
||||
|
||||
r.offers.SetState(id, StateCancelled)
|
||||
r.offers.Remove(id)
|
||||
r.spool.Remove(id)
|
||||
|
||||
offer.State = StateCancelled
|
||||
r.notifyWithdrawn(offer)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *receiver) receivedBytes(sender senderIdentity, id OfferID, index int) (int64, error) {
|
||||
offer, ok := r.offers.Get(sender.key, id)
|
||||
if !ok || index >= len(offer.Files) {
|
||||
return 0, ErrOfferNotFound
|
||||
}
|
||||
|
||||
received, err := r.spool.Received(id, index)
|
||||
if err != nil {
|
||||
log.Debugf("probe spool file: %v", err)
|
||||
return 0, fmt.Errorf("%w: read staged size", ErrStorage)
|
||||
}
|
||||
return received, nil
|
||||
}
|
||||
|
||||
func (r *receiver) upload(sender senderIdentity, id OfferID, index int, offset int64, body io.Reader) error {
|
||||
offer, ok := r.offers.Get(sender.key, id)
|
||||
if !ok || index >= len(offer.Files) {
|
||||
return ErrOfferNotFound
|
||||
}
|
||||
if offer.Decision != DecisionAccepted {
|
||||
return ErrNotAccepted
|
||||
}
|
||||
if offer.Files[index].Kind == KindText {
|
||||
return fmt.Errorf("%w: text payloads carry no body", ErrInvalidOffer)
|
||||
}
|
||||
|
||||
size := offer.Files[index].Size
|
||||
if offset < 0 || offset > size {
|
||||
return fmt.Errorf("%w: offset out of range", ErrInvalidOffer)
|
||||
}
|
||||
|
||||
received, err := r.spool.Write(id, index, offset, body, size)
|
||||
r.offers.SetProgress(id, index, received)
|
||||
r.notifyProgress(offer, index, received)
|
||||
|
||||
if err != nil {
|
||||
r.offers.SetState(id, StateFailed)
|
||||
r.notifyFailed(offer, err)
|
||||
log.Debugf("stage payload for offer %s file %d: %v", id, index, err)
|
||||
return fmt.Errorf("%w: stage payload", ErrStorage)
|
||||
}
|
||||
|
||||
if completed, ok := r.offers.Complete(id); ok {
|
||||
r.notifyCompleted(completed)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// expireOverdue reclaims offers past their decision deadline and stale spool data.
|
||||
func (r *receiver) expireOverdue() {
|
||||
for _, offer := range r.offers.ExpireOverdue() {
|
||||
r.spool.Remove(offer.ID)
|
||||
r.notifyFailed(offer, ErrExpired)
|
||||
}
|
||||
r.spool.Cleanup(r.spoolMaxAge, time.Now())
|
||||
}
|
||||
|
||||
func (r *receiver) close() {
|
||||
for _, offer := range r.offers.List() {
|
||||
r.offers.Remove(offer.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *receiver) notifyOffer(offer Offer) {
|
||||
if r.notifier != nil {
|
||||
r.notifier.OnOffer(offer)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *receiver) notifyProgress(offer Offer, index int, received int64) {
|
||||
if r.notifier != nil {
|
||||
r.notifier.OnProgress(offer, index, received)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *receiver) notifyCompleted(offer Offer) {
|
||||
if r.notifier != nil {
|
||||
r.notifier.OnCompleted(offer)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *receiver) notifyFailed(offer Offer, err error) {
|
||||
if r.notifier != nil {
|
||||
r.notifier.OnFailed(offer, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *receiver) notifyWithdrawn(offer Offer) {
|
||||
if r.notifier != nil {
|
||||
r.notifier.OnWithdrawn(offer)
|
||||
}
|
||||
}
|
||||
|
||||
func validateOffer(files []FileMeta) error {
|
||||
if len(files) == 0 {
|
||||
return fmt.Errorf("offer announces no files")
|
||||
}
|
||||
if len(files) > MaxOfferFiles {
|
||||
return fmt.Errorf("offer announces more than %d files", MaxOfferFiles)
|
||||
}
|
||||
|
||||
for _, f := range files {
|
||||
if !f.Kind.valid() {
|
||||
return fmt.Errorf("unknown payload kind %s", f.Kind)
|
||||
}
|
||||
if f.Kind == KindText {
|
||||
if len(f.Text) > MaxInlineTextSize {
|
||||
return fmt.Errorf("inline text exceeds %d bytes", MaxInlineTextSize)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if f.Size < 0 {
|
||||
return fmt.Errorf("negative file size")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
262
client/internal/filedrop/server.go
Normal file
262
client/internal/filedrop/server.go
Normal file
@@ -0,0 +1,262 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.zx2c4.com/wireguard/tun/netstack"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultSpoolMaxAge = 24 * time.Hour
|
||||
janitorInterval = 10 * time.Minute
|
||||
readHeaderTimeout = 30 * time.Second
|
||||
idleTimeout = 5 * time.Minute
|
||||
)
|
||||
|
||||
// PeerResolver maps the source overlay address of a connection to the peer that owns it.
|
||||
type PeerResolver interface {
|
||||
ResolvePeer(addr netip.Addr) (key PeerKey, name string, ok bool)
|
||||
}
|
||||
|
||||
// Notifier receives receiver-side transfer events for the platform layer to surface.
|
||||
type Notifier interface {
|
||||
OnOffer(offer Offer)
|
||||
OnProgress(offer Offer, index int, received int64)
|
||||
OnCompleted(offer Offer)
|
||||
OnFailed(offer Offer, err error)
|
||||
OnWithdrawn(offer Offer)
|
||||
}
|
||||
|
||||
// ServerConfig configures the receiving side.
|
||||
type ServerConfig struct {
|
||||
SpoolDir string
|
||||
Policy *PolicyStore
|
||||
Resolver PeerResolver
|
||||
Notifier Notifier
|
||||
OfferTTL time.Duration
|
||||
SpoolMaxAge time.Duration
|
||||
}
|
||||
|
||||
// Server serves the receiver over HTTP on the overlay address and owns the
|
||||
// listener and janitor lifecycle; the protocol logic itself lives in receiver.
|
||||
type Server struct {
|
||||
mu sync.RWMutex
|
||||
httpServer *http.Server
|
||||
listener net.Listener
|
||||
extraListeners []net.Listener
|
||||
netstackNet *netstack.Net
|
||||
|
||||
recv *receiver
|
||||
boundPort uint16
|
||||
|
||||
janitorStop context.CancelFunc
|
||||
janitorDone chan struct{}
|
||||
}
|
||||
|
||||
// NewServer builds a receiving server. It does not start listening.
|
||||
func NewServer(cfg ServerConfig) (*Server, error) {
|
||||
if cfg.Resolver == nil {
|
||||
return nil, errors.New("peer resolver is required")
|
||||
}
|
||||
if cfg.Policy == nil {
|
||||
return nil, errors.New("receiving policy is required")
|
||||
}
|
||||
|
||||
spool, err := NewSpool(cfg.SpoolDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create spool: %w", err)
|
||||
}
|
||||
|
||||
maxAge := cfg.SpoolMaxAge
|
||||
if maxAge <= 0 {
|
||||
maxAge = defaultSpoolMaxAge
|
||||
}
|
||||
|
||||
return &Server{recv: newReceiver(cfg, spool, maxAge)}, nil
|
||||
}
|
||||
|
||||
// SetNetstackNet routes listeners through the gVisor netstack instead of host sockets.
|
||||
func (s *Server) SetNetstackNet(n *netstack.Net) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.netstackNet = n
|
||||
}
|
||||
|
||||
// Offers returns the offer store, for the platform layer to accept, decline, and list.
|
||||
func (s *Server) Offers() *OfferStore {
|
||||
return s.recv.offers
|
||||
}
|
||||
|
||||
// Spool returns the staging area, so the platform layer can deliver completed payloads.
|
||||
func (s *Server) Spool() *Spool {
|
||||
return s.recv.spool
|
||||
}
|
||||
|
||||
// Policy returns the active profile's receiving policy store.
|
||||
func (s *Server) Policy() *PolicyStore {
|
||||
return s.recv.policy
|
||||
}
|
||||
|
||||
// BoundPort returns the port the server actually listens on, 0 when stopped.
|
||||
func (s *Server) BoundPort() uint16 {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.boundPort
|
||||
}
|
||||
|
||||
// Start binds the service to addr and serves until Stop.
|
||||
func (s *Server) Start(ctx context.Context, addr netip.AddrPort) error {
|
||||
s.mu.Lock()
|
||||
if s.httpServer != nil {
|
||||
s.mu.Unlock()
|
||||
return errors.New("file drop server is already running")
|
||||
}
|
||||
|
||||
ln, desc, err := s.createListener(ctx, addr)
|
||||
if err != nil && addr.Port() != 0 {
|
||||
log.Warnf("file drop port %d is unavailable, falling back to a dynamic port: %v", addr.Port(), err)
|
||||
ln, desc, err = s.createListener(ctx, netip.AddrPortFrom(addr.Addr(), 0))
|
||||
}
|
||||
if err != nil {
|
||||
s.mu.Unlock()
|
||||
return fmt.Errorf("create listener: %w", err)
|
||||
}
|
||||
|
||||
transport := &httpTransport{recv: s.recv}
|
||||
httpServer := &http.Server{
|
||||
Handler: transport.routes(),
|
||||
ReadHeaderTimeout: readHeaderTimeout,
|
||||
IdleTimeout: idleTimeout,
|
||||
}
|
||||
|
||||
janitorCtx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
|
||||
s.listener = ln
|
||||
s.httpServer = httpServer
|
||||
s.boundPort = listenerPort(ln, addr.Port())
|
||||
s.janitorStop = cancel
|
||||
s.janitorDone = done
|
||||
s.mu.Unlock()
|
||||
|
||||
go s.runJanitor(janitorCtx, done)
|
||||
go s.serve(httpServer, ln, desc)
|
||||
|
||||
log.Infof("file drop server started on %s", desc)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddListener serves the running service on an additional address, such as IPv6.
|
||||
func (s *Server) AddListener(ctx context.Context, addr netip.AddrPort) error {
|
||||
s.mu.Lock()
|
||||
httpServer := s.httpServer
|
||||
if httpServer == nil {
|
||||
s.mu.Unlock()
|
||||
return errors.New("file drop server is not running")
|
||||
}
|
||||
|
||||
ln, desc, err := s.createListener(ctx, addr)
|
||||
if err != nil {
|
||||
s.mu.Unlock()
|
||||
return fmt.Errorf("create listener: %w", err)
|
||||
}
|
||||
s.extraListeners = append(s.extraListeners, ln)
|
||||
s.mu.Unlock()
|
||||
|
||||
go s.serve(httpServer, ln, desc)
|
||||
|
||||
log.Infof("file drop server also listening on %s", desc)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop shuts the service down and releases the offers it was tracking. It is idempotent.
|
||||
func (s *Server) Stop() error {
|
||||
s.mu.Lock()
|
||||
httpServer := s.httpServer
|
||||
if httpServer == nil {
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
s.httpServer = nil
|
||||
s.listener = nil
|
||||
s.boundPort = 0
|
||||
extra := s.extraListeners
|
||||
s.extraListeners = nil
|
||||
stopJanitor, janitorDone := s.janitorStop, s.janitorDone
|
||||
s.janitorStop, s.janitorDone = nil, nil
|
||||
s.mu.Unlock()
|
||||
|
||||
if stopJanitor != nil {
|
||||
stopJanitor()
|
||||
<-janitorDone
|
||||
}
|
||||
|
||||
err := httpServer.Close()
|
||||
|
||||
for _, ln := range extra {
|
||||
if cerr := ln.Close(); cerr != nil {
|
||||
log.Debugf("close extra file drop listener: %v", cerr)
|
||||
}
|
||||
}
|
||||
|
||||
s.recv.close()
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("close: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) serve(httpServer *http.Server, ln net.Listener, desc string) {
|
||||
if err := httpServer.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
log.Errorf("file drop server error on %s: %v", desc, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) createListener(ctx context.Context, addr netip.AddrPort) (net.Listener, string, error) {
|
||||
if s.netstackNet != nil {
|
||||
ln, err := s.netstackNet.ListenTCPAddrPort(addr)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("listen on netstack: %w", err)
|
||||
}
|
||||
return ln, fmt.Sprintf("netstack %s", addr), nil
|
||||
}
|
||||
|
||||
var lc net.ListenConfig
|
||||
ln, err := lc.Listen(ctx, "tcp", net.TCPAddrFromAddrPort(addr).String())
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("listen: %w", err)
|
||||
}
|
||||
return ln, addr.String(), nil
|
||||
}
|
||||
|
||||
func (s *Server) runJanitor(ctx context.Context, done chan struct{}) {
|
||||
defer close(done)
|
||||
|
||||
ticker := time.NewTicker(janitorInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.recv.expireOverdue()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func listenerPort(ln net.Listener, requested uint16) uint16 {
|
||||
if tcpAddr, ok := ln.Addr().(*net.TCPAddr); ok {
|
||||
return uint16(tcpAddr.Port)
|
||||
}
|
||||
return requested
|
||||
}
|
||||
131
client/internal/filedrop/spool.go
Normal file
131
client/internal/filedrop/spool.go
Normal file
@@ -0,0 +1,131 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// Spool stages incoming payloads in an app-private directory.
|
||||
type Spool struct {
|
||||
root string
|
||||
}
|
||||
|
||||
// NewSpool prepares the spool directory tree under root.
|
||||
func NewSpool(root string) (*Spool, error) {
|
||||
if root == "" {
|
||||
return nil, fmt.Errorf("empty spool root")
|
||||
}
|
||||
if err := os.MkdirAll(root, 0o700); err != nil {
|
||||
return nil, fmt.Errorf("create spool root: %w", err)
|
||||
}
|
||||
return &Spool{root: root}, nil
|
||||
}
|
||||
|
||||
// Root returns the spool base directory.
|
||||
func (s *Spool) Root() string {
|
||||
return s.root
|
||||
}
|
||||
|
||||
// OfferDir returns the directory holding one offer's payloads.
|
||||
func (s *Spool) OfferDir(id OfferID) string {
|
||||
return filepath.Join(s.root, string(id))
|
||||
}
|
||||
|
||||
func (s *Spool) filePath(id OfferID, index int) string {
|
||||
return filepath.Join(s.OfferDir(id), strconv.Itoa(index))
|
||||
}
|
||||
|
||||
// Prepare creates the directory for an offer.
|
||||
func (s *Spool) Prepare(id OfferID) error {
|
||||
if err := os.MkdirAll(s.OfferDir(id), 0o700); err != nil {
|
||||
return fmt.Errorf("create offer dir: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Received returns how many bytes of one item are already staged.
|
||||
func (s *Spool) Received(id OfferID, index int) (int64, error) {
|
||||
info, err := os.Stat(s.filePath(id, index))
|
||||
if os.IsNotExist(err) {
|
||||
return 0, nil
|
||||
}
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("stat spool file: %w", err)
|
||||
}
|
||||
return info.Size(), nil
|
||||
}
|
||||
|
||||
// Write appends the payload at offset, truncating any bytes past it first.
|
||||
func (s *Spool) Write(id OfferID, index int, offset int64, r io.Reader, limit int64) (int64, error) {
|
||||
if offset < 0 {
|
||||
return 0, fmt.Errorf("negative offset %d", offset)
|
||||
}
|
||||
|
||||
path := s.filePath(id, index)
|
||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY, 0o600)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("open spool file: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := f.Close(); err != nil {
|
||||
log.Debugf("close spool file: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := f.Truncate(offset); err != nil {
|
||||
return 0, fmt.Errorf("truncate spool file: %w", err)
|
||||
}
|
||||
if _, err := f.Seek(offset, io.SeekStart); err != nil {
|
||||
return 0, fmt.Errorf("seek spool file: %w", err)
|
||||
}
|
||||
|
||||
written, err := io.Copy(f, io.LimitReader(r, limit-offset))
|
||||
if err != nil {
|
||||
return offset + written, fmt.Errorf("write spool file: %w", err)
|
||||
}
|
||||
return offset + written, nil
|
||||
}
|
||||
|
||||
// Path returns the staged path of one item for the platform layer to deliver from.
|
||||
func (s *Spool) Path(id OfferID, index int) string {
|
||||
return s.filePath(id, index)
|
||||
}
|
||||
|
||||
// Remove deletes an offer's staged payloads.
|
||||
func (s *Spool) Remove(id OfferID) {
|
||||
if err := os.RemoveAll(s.OfferDir(id)); err != nil {
|
||||
log.Debugf("remove spool dir: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup removes offer directories older than maxAge.
|
||||
func (s *Spool) Cleanup(maxAge time.Duration, now time.Time) {
|
||||
entries, err := os.ReadDir(s.root)
|
||||
if err != nil {
|
||||
log.Debugf("read spool root: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
log.Debugf("stat spool entry: %v", err)
|
||||
continue
|
||||
}
|
||||
if now.Sub(info.ModTime()) < maxAge {
|
||||
continue
|
||||
}
|
||||
if err := os.RemoveAll(filepath.Join(s.root, entry.Name())); err != nil {
|
||||
log.Debugf("remove stale spool dir: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
58
client/internal/filedrop/store.go
Normal file
58
client/internal/filedrop/store.go
Normal file
@@ -0,0 +1,58 @@
|
||||
package filedrop
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
)
|
||||
|
||||
const (
|
||||
namespacePolicy = "filedrop"
|
||||
namespaceHistory = "filedrop-history"
|
||||
)
|
||||
|
||||
// Store persists one profile's file drop state in namespaced sections.
|
||||
type Store interface {
|
||||
Get(namespace string, v any) (bool, error)
|
||||
Put(namespace string, v any) error
|
||||
}
|
||||
|
||||
type profileStore struct {
|
||||
prefs *profilemanager.Prefs
|
||||
}
|
||||
|
||||
// NewProfileStore returns the store backed by the profile's preferences.
|
||||
func NewProfileStore(prefs *profilemanager.Prefs) Store {
|
||||
if prefs == nil {
|
||||
return nil
|
||||
}
|
||||
return &profileStore{prefs: prefs}
|
||||
}
|
||||
|
||||
func (s *profileStore) Get(namespace string, v any) (bool, error) {
|
||||
return s.prefs.Get(namespace, v)
|
||||
}
|
||||
|
||||
func (s *profileStore) Put(namespace string, v any) error {
|
||||
return s.prefs.Put(namespace, v)
|
||||
}
|
||||
|
||||
func loadSection(store Store, namespace string, v any) error {
|
||||
if store == nil {
|
||||
return nil
|
||||
}
|
||||
if _, err := store.Get(namespace, v); err != nil {
|
||||
return fmt.Errorf("load %s: %w", namespace, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func saveSection(store Store, namespace string, v any) error {
|
||||
if store == nil {
|
||||
return nil
|
||||
}
|
||||
if err := store.Put(namespace, v); err != nil {
|
||||
return fmt.Errorf("save %s: %w", namespace, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,6 +1,8 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/pion/ice/v4"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
@@ -12,6 +14,7 @@ import (
|
||||
type Signaler struct {
|
||||
signal signal.Client
|
||||
wgPrivateKey wgtypes.Key
|
||||
filedropPort atomic.Uint32
|
||||
}
|
||||
|
||||
func NewSignaler(signal signal.Client, wgPrivateKey wgtypes.Key) *Signaler {
|
||||
@@ -44,6 +47,12 @@ func (s *Signaler) Ready() bool {
|
||||
return s.signal.Ready()
|
||||
}
|
||||
|
||||
// SetFiledropPort sets the file drop listen port advertised in offers and answers;
|
||||
// 0 means the well-known default and is not put on the wire.
|
||||
func (s *Signaler) SetFiledropPort(port uint16) {
|
||||
s.filedropPort.Store(uint32(port))
|
||||
}
|
||||
|
||||
// SignalOfferAnswer signals either an offer or an answer to remote peer
|
||||
func (s *Signaler) signalOfferAnswer(offerAnswer OfferAnswer, remoteKey string, bodyType sProto.Body_Type) error {
|
||||
var sessionIDBytes []byte
|
||||
@@ -57,6 +66,7 @@ func (s *Signaler) signalOfferAnswer(offerAnswer OfferAnswer, remoteKey string,
|
||||
msg, err := signal.MarshalCredential(s.wgPrivateKey, remoteKey, signal.CredentialPayload{
|
||||
Type: bodyType,
|
||||
WgListenPort: offerAnswer.WgListenPort,
|
||||
FiledropPort: uint16(s.filedropPort.Load()),
|
||||
Credential: &signal.Credential{
|
||||
UFrag: offerAnswer.IceCredentials.UFrag,
|
||||
Pwd: offerAnswer.IceCredentials.Pwd,
|
||||
|
||||
Reference in New Issue
Block a user