From 968ec6aacf79f7e49a2f251dc3bf90abfee914fc Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:41:25 +0300 Subject: [PATCH 01/12] feat(config): v1.19 configurability (network, reconnect, discovery, cache, pairing) - TCP plumbing: listener and identity advertise cfg.TCPPort; reconnect\n falls back to cfg.TCPPort when LastPort unknown (UDP discovery stays 1716)\n- [network] dial_timeout/handshake_timeout/sidechannel_timeout; shared\n transport.DialSidechannel replaces 4 per-plugin dialers (setup-only\n timeout, payload streaming unbounded, pin verification preserved)\n- [reconnect] initial_backoff/max_backoff/flap_threshold (overflow-safe)\n- [discovery] broadcast_interval/broadcast_idle_interval\n- [pairing] intent_ttl + listen_timeout (CLI client deadline = listen\n timeout + 10s headroom, overflow-clamped)\n- [cache] sms_attachments_dir/album_art_dir/contacts_dir overrides\n- [notifications] app_name reserved key; Filters() excludes it\n- config.Duration helper; Load validates positivity and ordering\n- tests across config/daemon/device/discovery/ipc/cmd; docs + example --- cmd/kcd/configurability_test.go | 18 ++ cmd/kcd/main.go | 14 +- docs/CLI.md | 39 +++ internal/config/config.go | 35 ++- internal/config/config_test.go | 226 ++++++++++++++++++ internal/config/plugins.go | 9 +- internal/config/runtime.go | 73 ++++++ internal/daemon/configurability_test.go | 46 ++++ internal/daemon/daemon.go | 13 +- internal/daemon/ipc_routes.go | 6 +- internal/daemon/plugins.go | 26 +- internal/daemon/transport.go | 57 ++--- internal/device/configurability_test.go | 38 +++ internal/device/device.go | 8 +- internal/device/registry.go | 26 +- internal/discovery/discovery.go | 15 +- internal/discovery/discovery_test.go | 10 + internal/ipc/configurability_test.go | 32 +++ internal/ipc/handler.go | 31 ++- internal/plugins/battery/battery.go | 22 +- internal/plugins/clipboard/clipboard.go | 34 +-- internal/plugins/contacts/contacts.go | 5 +- internal/plugins/mpris/artcache.go | 5 +- internal/plugins/mpris/mpris.go | 4 +- internal/plugins/notification/notification.go | 36 ++- internal/plugins/ping/ping.go | 9 +- internal/plugins/share/share.go | 11 +- internal/plugins/share/transfer.go | 33 +-- internal/plugins/sms/sms.go | 66 ++--- internal/plugins/telephony/telephony.go | 19 +- internal/transport/sidechannel.go | 70 ++++++ internal/transport/sidechannel_test.go | 183 ++++++++++++++ packaging/kcd.example.toml | 43 +++- pkg/client/client.go | 15 +- 34 files changed, 1061 insertions(+), 216 deletions(-) create mode 100644 cmd/kcd/configurability_test.go create mode 100644 internal/config/config_test.go create mode 100644 internal/config/runtime.go create mode 100644 internal/daemon/configurability_test.go create mode 100644 internal/device/configurability_test.go create mode 100644 internal/ipc/configurability_test.go create mode 100644 internal/transport/sidechannel.go create mode 100644 internal/transport/sidechannel_test.go diff --git a/cmd/kcd/configurability_test.go b/cmd/kcd/configurability_test.go new file mode 100644 index 0000000..2ec6b1d --- /dev/null +++ b/cmd/kcd/configurability_test.go @@ -0,0 +1,18 @@ +package main + +import ( + "testing" + "time" +) + +func TestPairListenDeadline(t *testing.T) { + for _, tc := range []struct{ input, want time.Duration }{ + {time.Minute, 70 * time.Second}, + {3 * time.Minute, 190 * time.Second}, + {time.Duration(1<<63 - 1), time.Duration(1<<63 - 1)}, + } { + if got := pairListenDeadline(tc.input); got != tc.want { + t.Errorf("%v: got %v, want %v", tc.input, got, tc.want) + } + } +} diff --git a/cmd/kcd/main.go b/cmd/kcd/main.go index d3ccb2d..54a7219 100644 --- a/cmd/kcd/main.go +++ b/cmd/kcd/main.go @@ -31,11 +31,21 @@ func getClient(c *cli.Context) (*client.Client, error) { return nil, fmt.Errorf("failed to load config: %w", err) } return &client.Client{ - SocketPath: cfg.SocketPath, - Timeout: 5 * time.Second, + SocketPath: cfg.SocketPath, + Timeout: 5 * time.Second, + PairListenTimeout: pairListenDeadline(config.Duration(cfg.Pairing.ListenTimeout)), }, nil } +// pairListenDeadline adds response overhead without overflowing a duration. +func pairListenDeadline(timeout time.Duration) time.Duration { + const maxDuration = time.Duration(1<<63 - 1) + if timeout > maxDuration-10*time.Second { + return maxDuration + } + return timeout + 10*time.Second +} + func main() { daemon.Version = version diff --git a/docs/CLI.md b/docs/CLI.md index 12a1efd..151a295 100644 --- a/docs/CLI.md +++ b/docs/CLI.md @@ -19,6 +19,45 @@ These flags apply to every command: --- +## Configuration + +Edit `$XDG_CONFIG_HOME/kcd/kcd.toml` (default `~/.config/kcd/kcd.toml`). +All settings are optional; see the annotated example in the packaging directory. +Apply changes with `systemctl --user restart kcd` (or restart `kcd daemon` +when running it directly). Timing, storage and notification branding settings +require a restart; reloading notification filters alone does not apply them. + +| Section | Settings and defaults | +|---|---| +| `[network]` | `dial_timeout = "5s"`, `handshake_timeout = "10s"`, `sidechannel_timeout = "15s"` | +| `[reconnect]` | `initial_backoff = "2s"`, `max_backoff = "5m"`, `flap_threshold = "15s"` | +| `[discovery]` | `broadcast_interval = "30s"`, `broadcast_idle_interval = "60s"` | +| `[pairing]` | `intent_ttl = "5m"`, `listen_timeout = "60s"`; existing `timeout_secs = 30` still controls the pairing response wait | +| `[cache]` | `sms_attachments_dir = ""`, `album_art_dir = ""`, `contacts_dir = ""` | +| `[notifications]` | `app_name = "KDE Connect"`; per-app `"show"`/`"silent"` filters and `"*"` fallback remain supported | +| `[ping]` | `app_name = ""` inherits the notification application name; a non-empty value overrides it | + +Durations use Go syntax such as `"750ms"`, `"10s"`, or `"1m30s"` and must be +strictly positive. Maximum reconnect backoff must be at least the initial +backoff; idle broadcast interval must be at least the normal interval. Invalid +values are rejected when loading configuration, including in the CLI. Discovery +intervals only apply while on-demand broadcast is active, not during connected +steady state; mDNS advertisement is unchanged. + +Empty cache overrides preserve existing paths: SMS attachments use the system +temporary directory's `kcd/sms-attachments`, album art uses +`$XDG_CACHE_HOME/kcd/art` (normally `~/.cache/kcd/art`), and contacts use +`$XDG_DATA_HOME/kcd/contacts` (normally `~/.local/share/kcd/contacts`). Use absolute +paths for overrides. Changing directories does not migrate existing files. + +`notifications.app_name` is reserved branding metadata, never a per-app filter. +An explicit `ping.app_name = "KDE Connect"` in an older configuration remains an +override even after changing the global name; remove it or set it to `""` to +inherit. Protocol version, payload limits, packet buffers, queues and TCP +keepalive remain fixed implementation settings, not configuration knobs. + +--- + ## daemon Start the `kcd` background daemon. This is the only command that does not connect to a running daemon — it *is* the daemon. diff --git a/internal/config/config.go b/internal/config/config.go index dbf2da8..4de20cc 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -27,6 +27,10 @@ type Config struct { LogLevel string `toml:"log_level"` // "debug", "info", "warn", "error" (or "quiet") // AutoAcceptPairing was removed in favor of `kcd pair` (listen mode). // Old config values are silently ignored by the TOML parser. + Network NetworkConfig `toml:"network"` + Reconnect ReconnectConfig `toml:"reconnect"` + Discovery DiscoveryConfig `toml:"discovery"` + Cache CacheConfig `toml:"cache"` Plugins PluginConfig `toml:"plugins"` Commands map[string]string `toml:"commands"` CommandsPerDevice map[string]map[string]string `toml:"commands_per_device"` @@ -65,6 +69,9 @@ func Defaults() *Config { c.TCPPort = 1716 c.LogLevel = "info" + c.Network = NetworkConfig{DialTimeout: "5s", HandshakeTimeout: "10s", SidechannelTimeout: "15s"} + c.Reconnect = ReconnectConfig{InitialBackoff: "2s", MaxBackoff: "5m", FlapThreshold: "15s"} + c.Discovery = DiscoveryConfig{BroadcastInterval: "30s", BroadcastIdleInterval: "60s"} c.Plugins.Defaults() c.Commands = make(map[string]string) c.CommandsPerDevice = make(map[string]map[string]string) @@ -100,11 +107,17 @@ func Load(path string) (*Config, error) { } cfg.ConfigPath = path + if err := cfg.Validate(); err != nil { + return nil, err + } return cfg, nil } // Validate checks required fields and returns an error if any are invalid. func (c *Config) Validate() error { + if err := c.validateDurations(); err != nil { + return err + } if c.DeviceName == "" { return fmt.Errorf("config: device_name is required") } @@ -235,9 +248,29 @@ func configPath(filename string, isRuntime bool) string { return filepath.Join(dir, filename) } -// NotificationConfig controls per-app notification filtering. +// NotificationConfig controls notification branding and per-app filtering. +// The reserved app_name key is not a filter. type NotificationConfig map[string]string +// AppName returns the desktop notification application name. +func (c NotificationConfig) AppName() string { + if name := c["app_name"]; name != "" { + return name + } + return "KDE Connect" +} + +// Filters returns an independent copy containing only per-app filter entries. +func (c NotificationConfig) Filters() NotificationConfig { + filters := make(NotificationConfig, len(c)) + for app, action := range c { + if app != "app_name" { + filters[app] = action + } + } + return filters +} + // generateDeviceID produces a UUIDv4 with dashes replaced by underscores. func generateDeviceID() (string, error) { var uuid [16]byte diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..eb024da --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,226 @@ +package config + +import ( + "os" + "path/filepath" + "reflect" + "strings" + "testing" + "time" + + "github.com/BurntSushi/toml" +) + +func TestDefaults(t *testing.T) { + cfg := Defaults() + if err := cfg.Validate(); err != nil { + t.Fatal(err) + } + if cfg.Network != (NetworkConfig{"5s", "10s", "15s"}) { + t.Errorf("network defaults: %+v", cfg.Network) + } + if cfg.Reconnect != (ReconnectConfig{"2s", "5m", "15s"}) { + t.Errorf("reconnect defaults: %+v", cfg.Reconnect) + } + if cfg.Discovery != (DiscoveryConfig{"30s", "60s"}) { + t.Errorf("discovery defaults: %+v", cfg.Discovery) + } + if cfg.Pairing.IntentTTL != "5m" || cfg.Pairing.ListenTimeout != "60s" || cfg.Pairing.TimeoutSecs != 30 { + t.Errorf("pairing defaults: %+v", cfg.Pairing) + } + if cfg.Cache != (CacheConfig{}) || cfg.Ping.AppName != "" || cfg.Notifications.AppName() != "KDE Connect" { + t.Fatal("default cache paths or notification inheritance changed") + } +} + +func loadTOML(t *testing.T, text string) (*Config, error) { + t.Helper() + path := filepath.Join(t.TempDir(), "kcd.toml") + if err := os.WriteFile(path, []byte(text), 0600); err != nil { + t.Fatal(err) + } + return Load(path) +} + +func TestLoadDefaults(t *testing.T) { + for _, contents := range []string{"", "device_name = 'legacy'\n[plugins]\nping = false\n"} { + cfg, err := loadTOML(t, contents) + if err != nil { + t.Fatal(err) + } + if cfg.Network != Defaults().Network || cfg.Pairing.IntentTTL != "5m" || cfg.ConfigPath == "" { + t.Fatal("omitted fields did not retain defaults") + } + } + path := filepath.Join(t.TempDir(), "absent.toml") + cfg, err := Load(path) + if err != nil || cfg.Validate() != nil { + t.Fatalf("missing file must return valid defaults: %v", err) + } + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Fatal("Load unexpectedly created missing file") + } +} + +func TestLoadOverridesAndRoundTrip(t *testing.T) { + cfg, err := loadTOML(t, ` +[network] +dial_timeout = "750ms" +handshake_timeout = "12s" +sidechannel_timeout = "25s" +[reconnect] +initial_backoff = "3s" +max_backoff = "6m" +flap_threshold = "20s" +[discovery] +broadcast_interval = "45s" +broadcast_idle_interval = "90s" +[cache] +sms_attachments_dir = "/tmp/sms" +album_art_dir = "/tmp/art" +contacts_dir = "/tmp/contacts" +[pairing] +intent_ttl = "7m" +listen_timeout = "2m" +[notifications] +app_name = "My daemon" +"*" = "show" +"com.example" = "silent" +[ping] +app_name = "KDE Connect" +`) + if err != nil { + t.Fatal(err) + } + if cfg.Network != (NetworkConfig{"750ms", "12s", "25s"}) || cfg.Reconnect != (ReconnectConfig{"3s", "6m", "20s"}) || cfg.Discovery != (DiscoveryConfig{"45s", "90s"}) { + t.Fatal("duration overrides not decoded") + } + if cfg.Cache != (CacheConfig{"/tmp/sms", "/tmp/art", "/tmp/contacts"}) || cfg.Pairing.IntentTTL != "7m" || cfg.Pairing.ListenTimeout != "2m" { + t.Fatal("cache/pairing overrides not decoded") + } + if cfg.Notifications.AppName() != "My daemon" || cfg.Ping.AppName != "KDE Connect" { + t.Fatal("explicit ping override must survive, including the historical default name") + } + if !reflect.DeepEqual(cfg.Notifications.Filters(), NotificationConfig{"*": "show", "com.example": "silent"}) { + t.Fatal("notification filter decode failed") + } + path := filepath.Join(t.TempDir(), "saved.toml") + if err := cfg.Save(path); err != nil { + t.Fatal(err) + } + reloaded, err := Load(path) + if err != nil { + t.Fatal(err) + } + cfg.ConfigPath = path + if !reflect.DeepEqual(cfg, reloaded) { + t.Fatal("config changed after save/load round trip") + } +} + +func TestDurationValidation(t *testing.T) { + fields := []struct{ section, key string }{ + {"network", "dial_timeout"}, {"network", "handshake_timeout"}, {"network", "sidechannel_timeout"}, + {"reconnect", "initial_backoff"}, {"reconnect", "max_backoff"}, {"reconnect", "flap_threshold"}, + {"discovery", "broadcast_interval"}, {"discovery", "broadcast_idle_interval"}, + {"pairing", "intent_ttl"}, {"pairing", "listen_timeout"}, + } + for _, field := range fields { + for _, value := range []string{"", "nonsense", "10", "0", "0s", "-1s", "9999999999999999999h"} { + t.Run(field.section+"."+field.key+"/"+value, func(t *testing.T) { + text := "[" + field.section + "]\n" + field.key + " = '" + value + "'\n" + if _, err := loadTOML(t, text); err == nil || !strings.Contains(err.Error(), field.section+"."+field.key) { + t.Fatalf("Load error = %v; want field-specific validation error", err) + } + }) + } + } +} + +func TestDurationRelationships(t *testing.T) { + for _, section := range []struct{ name, initial, maximum string }{ + {"reconnect", "initial_backoff", "max_backoff"}, + {"discovery", "broadcast_interval", "broadcast_idle_interval"}, + } { + for _, maximum := range []string{"999ms", "1s", "2s"} { + text := "[" + section.name + "]\n" + section.initial + " = '1s'\n" + section.maximum + " = '" + maximum + "'\n" + _, err := loadTOML(t, text) + if maximum == "999ms" { + if err == nil || !strings.Contains(err.Error(), section.name+"."+section.maximum) { + t.Fatalf("expected relationship error, got %v", err) + } + } else if err != nil { + t.Fatalf("equal or larger maximum should be valid: %v", err) + } + } + } +} + +func TestNotificationConfig(t *testing.T) { + for _, cfg := range []NotificationConfig{nil, {}, {"app_name": ""}, {"*": "silent"}} { + if cfg.AppName() != "KDE Connect" { + t.Fatal("missing or empty app_name must use default") + } + filters := cfg.Filters() + if _, ok := filters["app_name"]; ok { + t.Fatal("reserved app_name leaked into filters") + } + filters["new"] = "show" + if _, ok := cfg["new"]; ok { + t.Fatal("Filters returned shared map") + } + } + cfg := NotificationConfig{"app_name": "Custom", "*": "silent", "app": "show"} + filters := cfg.Filters() + filters["app"] = "silent" + cfg["*"] = "show" + if cfg.AppName() != "Custom" || cfg["app"] != "show" || filters["*"] != "silent" { + t.Fatal("filter map copy or app name failed") + } +} + +func TestDuration(t *testing.T) { + for value, want := range map[string]time.Duration{"1ns": time.Nanosecond, "750ms": 750 * time.Millisecond, "1m30s": 90 * time.Second} { + if got := Duration(value); got != want { + t.Errorf("Duration(%q) = %v, want %v", value, got, want) + } + } + t.Run("unvalidated value", func(t *testing.T) { + defer func() { + if recover() == nil { + t.Fatal("malformed duration should expose programming error") + } + }() + Duration("not-a-duration") + }) +} + +func TestLoadErrors(t *testing.T) { + for _, text := range []string{"[", "[network]\ndial_timeout = 2", "tcp_port = 0"} { + if _, err := loadTOML(t, text); err == nil { + t.Fatalf("expected error loading %q", text) + } + } + if _, err := Load(t.TempDir()); err == nil { + t.Fatal("expected read error for directory") + } +} + +func TestExampleConfig(t *testing.T) { + path := filepath.Join("..", "..", "packaging", "kcd.example.toml") + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(cfg.Plugins, Defaults().Plugins) { + t.Fatal("example changed plugin defaults") + } + var parsed Config + metadata, err := toml.DecodeFile(path, &parsed) + if err != nil { + t.Fatal(err) + } + if unknown := metadata.Undecoded(); len(unknown) != 0 { + t.Fatalf("unknown/misplaced example settings: %v", unknown) + } +} diff --git a/internal/config/plugins.go b/internal/config/plugins.go index 994b5b6..71cc8cc 100644 --- a/internal/config/plugins.go +++ b/internal/config/plugins.go @@ -85,7 +85,9 @@ type PingConfig struct { } type PairingConfig struct { - TimeoutSecs int `toml:"timeout_secs"` + TimeoutSecs int `toml:"timeout_secs"` + IntentTTL string `toml:"intent_ttl"` + ListenTimeout string `toml:"listen_timeout"` } type MousepadConfig struct { @@ -165,12 +167,15 @@ func (c *SFTPConfig) Defaults() { } func (c *PingConfig) Defaults() { - c.AppName = "KDE Connect" + // Empty inherits Notifications.AppName(); a non-empty value overrides it. + c.AppName = "" c.Icon = "smartphone" } func (c *PairingConfig) Defaults() { c.TimeoutSecs = 30 + c.IntentTTL = "5m" + c.ListenTimeout = "60s" } func (c *MousepadConfig) Defaults() { diff --git a/internal/config/runtime.go b/internal/config/runtime.go new file mode 100644 index 0000000..acff8d2 --- /dev/null +++ b/internal/config/runtime.go @@ -0,0 +1,73 @@ +package config + +import ( + "fmt" + "time" +) + +// NetworkConfig controls connection setup deadlines, not transfer size limits. +type NetworkConfig struct { + DialTimeout string `toml:"dial_timeout"` + HandshakeTimeout string `toml:"handshake_timeout"` + SidechannelTimeout string `toml:"sidechannel_timeout"` +} + +// ReconnectConfig controls retry delays and the minimum stable connection age. +type ReconnectConfig struct { + InitialBackoff string `toml:"initial_backoff"` + MaxBackoff string `toml:"max_backoff"` + FlapThreshold string `toml:"flap_threshold"` +} + +// DiscoveryConfig controls intervals while on-demand UDP discovery is running. +type DiscoveryConfig struct { + BroadcastInterval string `toml:"broadcast_interval"` + BroadcastIdleInterval string `toml:"broadcast_idle_interval"` +} + +// CacheConfig overrides storage directories. Empty values retain plugin defaults. +type CacheConfig struct { + SMSAttachmentsDir string `toml:"sms_attachments_dir"` + AlbumArtDir string `toml:"album_art_dir"` + ContactsDir string `toml:"contacts_dir"` +} + +// Duration parses a duration from a validated Config. Call Load or Validate first. +// Invalid input is a programming error; validation supplies user-facing errors. +func Duration(value string) time.Duration { + d, err := time.ParseDuration(value) + if err != nil { + panic(fmt.Sprintf("config: unvalidated duration %q: %v", value, err)) + } + return d +} + +func (c *Config) validateDurations() error { + for _, setting := range []struct{ name, value string }{ + {"network.dial_timeout", c.Network.DialTimeout}, + {"network.handshake_timeout", c.Network.HandshakeTimeout}, + {"network.sidechannel_timeout", c.Network.SidechannelTimeout}, + {"reconnect.initial_backoff", c.Reconnect.InitialBackoff}, + {"reconnect.max_backoff", c.Reconnect.MaxBackoff}, + {"reconnect.flap_threshold", c.Reconnect.FlapThreshold}, + {"discovery.broadcast_interval", c.Discovery.BroadcastInterval}, + {"discovery.broadcast_idle_interval", c.Discovery.BroadcastIdleInterval}, + {"pairing.intent_ttl", c.Pairing.IntentTTL}, + {"pairing.listen_timeout", c.Pairing.ListenTimeout}, + } { + d, err := time.ParseDuration(setting.value) + if err != nil { + return fmt.Errorf("config: invalid %s %q: %w", setting.name, setting.value, err) + } + if d <= 0 { + return fmt.Errorf("config: %s must be greater than zero", setting.name) + } + } + if Duration(c.Reconnect.MaxBackoff) < Duration(c.Reconnect.InitialBackoff) { + return fmt.Errorf("config: reconnect.max_backoff must be >= reconnect.initial_backoff") + } + if Duration(c.Discovery.BroadcastIdleInterval) < Duration(c.Discovery.BroadcastInterval) { + return fmt.Errorf("config: discovery.broadcast_idle_interval must be >= discovery.broadcast_interval") + } + return nil +} diff --git a/internal/daemon/configurability_test.go b/internal/daemon/configurability_test.go new file mode 100644 index 0000000..75da598 --- /dev/null +++ b/internal/daemon/configurability_test.go @@ -0,0 +1,46 @@ +package daemon + +import ( + "context" + "encoding/json" + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" + "net" + "testing" +) + +func TestReconnectConfiguredPort(t *testing.T) { + dev := device.NewDevice("peer", "Peer", "phone", zap.NewNop()) + if got := reconnectPort(dev, 1816); got != 1816 { + t.Fatalf("fallback = %d", got) + } + dev.SetLastPort(1916) + if got := reconnectPort(dev, 1816); got != 1916 { + t.Fatalf("peer port = %d", got) + } +} + +// A canceled dial cannot touch the network. The test documents acceptance of +// a non-default TCP port without changing the discovery protocol port. +func TestDialConfiguredPortCanceled(t *testing.T) { + cfg := config.Defaults() + cfg.TCPPort = 1816 + pkt, err := protocol.NewIdentityPacket("local", "Local", "desktop", cfg.TCPPort, nil, nil) + if err != nil { + t.Fatal(err) + } + var body protocol.IdentityBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + t.Fatal(err) + } + if body.TCPPort != cfg.TCPPort { + t.Fatalf("identity port = %d", body.TCPPort) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + logger := zap.NewNop() + DialDevice(ctx, net.IPv4(127, 0, 0, 1), cfg.TCPPort, "peer", protocol.ProtocolVersion, pkt, nil, device.NewRegistry(nil), plugin.NewRegistry(logger), "local", logger, true, cfg) +} diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index 3649405..234c64b 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -159,12 +159,12 @@ func Run(ctx context.Context, cfg *config.Config) error { // 4. Broadcast Controller (starts stopped — activated by `kcd pair`) incomingCaps, outgoingCaps := plugins.Capabilities() - identity, err := protocol.NewIdentityPacket(cfg.DeviceID, cfg.DeviceName, "desktop", 1716, incomingCaps, outgoingCaps) + identity, err := protocol.NewIdentityPacket(cfg.DeviceID, cfg.DeviceName, "desktop", cfg.TCPPort, incomingCaps, outgoingCaps) if err != nil { return err } - bc := discovery.NewBroadcasterController(identity, 30*time.Second, logger, devices.AllPairedDevicesConnected) + bc := discovery.NewBroadcasterController(identity, config.Duration(cfg.Discovery.BroadcastInterval), logger, devices.AllPairedDevicesConnected, config.Duration(cfg.Discovery.BroadcastIdleInterval)) // mDNS advertisement is always on: unlike UDP broadcast it is // responder-only (zero idle timers), so phones keep a standing @@ -189,6 +189,7 @@ func Run(ctx context.Context, cfg *config.Config) error { // 5. IPC Server handler := ipc.NewHandler(devices, plugins, pairPlugin, statePath, bus, pruneThreshold) + handler.SetPairListenTimeout(config.Duration(cfg.Pairing.ListenTimeout)) registerIPCRoutes(handler, cfg, devices, plugins, bc, ctx, tlsCfg, logger, startedAt) @@ -204,7 +205,7 @@ func Run(ctx context.Context, cfg *config.Config) error { if dev.State() == device.StatePaired { return fmt.Errorf("device already paired") } - dev.RequestPairDial() + dev.RequestPairDial(config.Duration(cfg.Pairing.IntentTTL)) ip, port := dev.DiscoveryAddr() if ip == nil { // No address yet — the flagged dial fires on the next @@ -215,7 +216,7 @@ func Run(ctx context.Context, cfg *config.Config) error { port = 1716 } go func() { - DialDevice(ctx, ip, port, deviceID, protocol.ProtocolVersion, identity, tlsCfg, devices, plugins, cfg.DeviceID, logger, true) + DialDevice(ctx, ip, port, deviceID, protocol.ProtocolVersion, identity, tlsCfg, devices, plugins, cfg.DeviceID, logger, true, cfg) if !dev.IsConnected() { logger.Warn("on-demand pair dial failed", zap.String("device_id", deviceID)) return @@ -249,7 +250,7 @@ func Run(ctx context.Context, cfg *config.Config) error { }() // 6. Transport Layer - go runTransport(ctx, tlsCfg, bc, identity, devices, plugins, cfg.DeviceID, logger) + go runTransport(ctx, tlsCfg, bc, identity, devices, plugins, cfg.DeviceID, logger, cfg) // Wait for context cancellation (SIGTERM) if notifySocket := os.Getenv("NOTIFY_SOCKET"); notifySocket != "" { @@ -294,7 +295,7 @@ func Run(ctx context.Context, cfg *config.Config) error { // Reload notification filters. if pl, ok := plugins.GetByName("Notification"); ok { - pl.(*notification.NotificationPlugin).SetFilters(newCfg.Notifications) + pl.(*notification.NotificationPlugin).SetFilters(newCfg.Notifications.Filters()) logger.Info("reloaded notification filters") } diff --git a/internal/daemon/ipc_routes.go b/internal/daemon/ipc_routes.go index 71af460..581348b 100644 --- a/internal/daemon/ipc_routes.go +++ b/internal/daemon/ipc_routes.go @@ -22,7 +22,7 @@ import ( func registerIPCRoutes(handler *ipc.Handler, cfg *config.Config, devices *device.Registry, plugins *plugin.Registry, bc *discovery.BroadcasterController, ctx context.Context, tlsCfg *tls.Config, logger *zap.Logger, startedAt time.Time) { if cfg.Plugins.Notification { if notifPl, ok := plugins.GetByName("Notification"); ok { - notifPl.(*notification.NotificationPlugin).SetFilters(cfg.Notifications) + notifPl.(*notification.NotificationPlugin).SetFilters(cfg.Notifications.Filters()) } } @@ -68,11 +68,11 @@ func registerIPCRoutes(handler *ipc.Handler, cfg *config.Config, devices *device go func() { incomingCaps, outgoingCaps := plugins.Capabilities() - identityPkt, err := protocol.NewIdentityPacket(cfg.DeviceID, cfg.DeviceName, "desktop", 1716, incomingCaps, outgoingCaps) + identityPkt, err := protocol.NewIdentityPacket(cfg.DeviceID, cfg.DeviceName, "desktop", cfg.TCPPort, incomingCaps, outgoingCaps) if err != nil { return } - DialDevice(ctx, addr, 1716, "manual", protocol.ProtocolVersion, identityPkt, tlsCfg, devices, plugins, cfg.DeviceID, logger, true) + DialDevice(ctx, addr, 1716, "manual", protocol.ProtocolVersion, identityPkt, tlsCfg, devices, plugins, cfg.DeviceID, logger, true, cfg) }() return ipc.Response{OK: true} diff --git a/internal/daemon/plugins.go b/internal/daemon/plugins.go index 9eaedf9..6c46ba8 100644 --- a/internal/daemon/plugins.go +++ b/internal/daemon/plugins.go @@ -28,38 +28,42 @@ import ( "github.com/bethropolis/kcd/internal/plugins/sms" "github.com/bethropolis/kcd/internal/plugins/systemvolume" "github.com/bethropolis/kcd/internal/plugins/telephony" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) func setupPlugins(cfg *config.Config, bus *events.Bus, tlsCfg *tls.Config, logger *zap.Logger, devices *device.Registry, localCert *x509.Certificate, saveDevices func(), plugins *plugin.Registry) *pair.PairPlugin { + sidechannel := transport.SidechannelOptions{ + Timeout: config.Duration(cfg.Network.SidechannelTimeout), + } pairPlugin := pair.NewPairPlugin(devices, localCert, cfg.Pairing, saveDevices, bus, logger) plugins.Register(pairPlugin) if cfg.Plugins.Battery { - plugins.Register(battery.NewBatteryPlugin(cfg.Battery, bus, logger)) + plugins.Register(battery.NewBatteryPlugin(cfg.Battery, bus, logger, cfg.Notifications)) } if cfg.Plugins.Notification { - plugins.Register(notification.NewNotificationPlugin(cfg.Notification, bus, tlsCfg, logger)) + plugins.Register(notification.NewNotificationPlugin(cfg.Notification, bus, tlsCfg, logger, sidechannel)) } if cfg.Plugins.Clipboard { - plugins.Register(clipboard.NewClipboardPlugin(tlsCfg, logger, cfg.Clipboard.PushOnConnect)) + plugins.Register(clipboard.NewClipboardPlugin(tlsCfg, logger, cfg.Clipboard.PushOnConnect, sidechannel)) } if cfg.Plugins.Share { - plugins.Register(share.NewSharePlugin(cfg.DownloadDir, cfg.Share, tlsCfg, bus, logger)) + plugins.Register(share.NewSharePlugin(cfg.DownloadDir, cfg.Share, tlsCfg, bus, logger, sidechannel)) } if cfg.Plugins.RunCommand { plugins.Register(runcommand.NewRunCommandPlugin(cfg.Commands, cfg.CommandsPerDevice, logger)) } if cfg.Plugins.Ping { - plugins.Register(ping.NewPingPlugin(cfg.Ping, bus, logger)) + plugins.Register(ping.NewPingPlugin(cfg.Ping, bus, logger, cfg.Notifications)) } if cfg.Plugins.Telephony { - plugins.Register(telephony.NewTelephonyPlugin(bus, logger)) + plugins.Register(telephony.NewTelephonyPlugin(bus, logger, cfg.Notifications)) } if cfg.Plugins.Connectivity { plugins.Register(connectivity.NewConnectivityPlugin(bus)) } if cfg.Plugins.MPRIS { - plugins.Register(mpris.NewMPRISPlugin(tlsCfg, bus, cfg.Plugins.PauseMusic, logger)) + plugins.Register(mpris.NewMPRISPlugin(tlsCfg, bus, cfg.Plugins.PauseMusic, logger, cfg.Cache.AlbumArtDir)) } if cfg.Plugins.Mousepad { plugins.Register(mousepad.NewMousepadPlugin(cfg.Mousepad, logger)) @@ -83,10 +87,14 @@ func setupPlugins(cfg *config.Config, bus *events.Bus, tlsCfg *tls.Config, logge plugins.Register(systemvolume.NewSystemVolumePlugin(bus, logger)) } if cfg.Plugins.SMS { - plugins.Register(sms.NewSMSPlugin(cfg.SMS, bus, tlsCfg, logger)) + plugins.Register(sms.NewSMSPlugin(cfg.SMS, bus, tlsCfg, logger, sms.Options{ + CacheDir: cfg.Cache.SMSAttachmentsDir, + Sidechannel: sidechannel, + Notifications: cfg.Notifications, + })) } if cfg.Plugins.Contacts { - plugins.Register(contacts.NewContactsPlugin(bus, logger)) + plugins.Register(contacts.NewContactsPlugin(bus, logger, cfg.Cache.ContactsDir)) } if cfg.Plugins.RemoteSystemVolume { plugins.Register(remotesystemvolume.NewRemoteSystemVolumePlugin(bus, logger)) diff --git a/internal/daemon/transport.go b/internal/daemon/transport.go index 478f2e0..d525e30 100644 --- a/internal/daemon/transport.go +++ b/internal/daemon/transport.go @@ -10,6 +10,7 @@ import ( "time" "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/discovery" "github.com/bethropolis/kcd/internal/plugin" @@ -18,12 +19,6 @@ import ( "go.uber.org/zap" ) -// reconnectFlapThreshold is the minimum connection lifetime for it to count -// as genuinely stable. A connection that drops sooner is treated as a flap -// (e.g. a peer that keeps dying), so the auto-reconnect backoff continues -// escalating instead of resetting to the 2s floor on every drop. -const reconnectFlapThreshold = 15 * time.Second - // validDialPort reports whether a discovery-advertised TCP port is usable. // Port 0 and out-of-range values come from malformed or hostile packets and // must never reach the dialer (port 0 would dial ":0"). @@ -35,7 +30,7 @@ func validDialPort(port int) bool { // inside reconnectCooldown are skipped so post-roam bursts can't complete // near-simultaneously and churn the peer's duplicate resolution. // Explicit user actions (pair intent, manual connect) pass force=true. -func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID string, targetProto int, identity *protocol.Packet, cfg *tls.Config, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, force bool) { +func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID string, targetProto int, identity *protocol.Packet, cfg *tls.Config, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, force bool, opts *config.Config) { if targetIP == nil || !validDialPort(targetPort) { logger.Debug("refusing to dial invalid target", zap.String("device_id", targetID), @@ -55,7 +50,7 @@ func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID s logger.Debug("dialing discovered device", zap.String("device_id", targetID), zap.String("addr", addr)) dialer := &net.Dialer{ - Timeout: 5 * time.Second, + Timeout: config.Duration(opts.Network.DialTimeout), } conn, err := dialer.DialContext(ctx, "tcp", addr) if err != nil { @@ -99,7 +94,7 @@ func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID s // KDE Connect inverts TLS roles: TCP client acts as TLS server tlsConn := tls.Server(conn, cfg) - handshakeCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + handshakeCtx, cancel := context.WithTimeout(ctx, config.Duration(opts.Network.HandshakeTimeout)) defer cancel() if err := tlsConn.HandshakeContext(handshakeCtx); err != nil { tlsConn.Close() @@ -109,7 +104,7 @@ func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID s transConn := transport.NewConn(tlsConn) // Ensure the connection is closed if handleNewConnection fails mid-setup - if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger); err != nil { + if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger, opts); err != nil { logger.Debug("new connection setup failed", zap.Error(err)) transConn.Close() } @@ -138,9 +133,9 @@ func shouldEphemeralClose(dev *device.Device, pairingMode bool) bool { return true } -func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.BroadcasterController, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger) { +func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.BroadcasterController, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, opts *config.Config) { // TCP Listener - tcpListener, err := transport.Listen(ctx, ":1716") + tcpListener, err := transport.Listen(ctx, fmt.Sprintf(":%d", opts.TCPPort)) if err != nil { logger.Error("failed to start TCP listener", zap.Error(err)) return @@ -256,7 +251,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste if p := dev.LastPort(); validDialPort(p) { port = p } - go DialDevice(ctx, ip, port, body.DeviceID, body.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false) + go DialDevice(ctx, ip, port, body.DeviceID, body.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) } } else if dev.IsConnected() { if currIP := dev.RemoteIP(); currIP == nil || !currIP.Equal(ip) { @@ -267,7 +262,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste if p := dev.LastPort(); validDialPort(p) { port = p } - go DialDevice(ctx, ip, port, body.DeviceID, body.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false) + go DialDevice(ctx, ip, port, body.DeviceID, body.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) } } } @@ -287,7 +282,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste if pairingMode || dev.ConsumePairDial() { // Spawn goroutine to prevent blocking the discovery listener go func(targetIP net.IP, targetPort int, targetID string, targetProto int) { - DialDevice(ctx, targetIP, targetPort, targetID, targetProto, identity, cfg, devices, plugins, localDeviceID, logger, true) + DialDevice(ctx, targetIP, targetPort, targetID, targetProto, identity, cfg, devices, plugins, localDeviceID, logger, true, opts) }(ip, tcpPort, body.DeviceID, body.ProtocolVersion) return } @@ -299,7 +294,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste } // Spawn goroutine to prevent blocking the discovery listener go func(targetIP net.IP, targetPort int, targetID string, targetProto int) { - DialDevice(ctx, targetIP, targetPort, targetID, targetProto, identity, cfg, devices, plugins, localDeviceID, logger, false) + DialDevice(ctx, targetIP, targetPort, targetID, targetProto, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) }(ip, tcpPort, body.DeviceID, body.ProtocolVersion) } // Otherwise the device already had its ephemeral dial for this @@ -359,7 +354,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste protocol.ReleasePacket(preTlsPkt) tlsConn := tls.Client(newConn, cfg) - handshakeCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + handshakeCtx, cancel := context.WithTimeout(ctx, config.Duration(opts.Network.HandshakeTimeout)) defer cancel() if err := tlsConn.HandshakeContext(handshakeCtx); err != nil { return @@ -368,7 +363,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste transConn := transport.NewConn(tlsConn) c = nil // Prevent defer from closing the active connection - if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger); err != nil { + if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger, opts); err != nil { logger.Debug("new connection setup failed", zap.Error(err)) transConn.Close() } @@ -380,7 +375,7 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste } // Returning an error ensures the caller can close the connection if it fails mid-setup. -func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, cfg *tls.Config, logger *zap.Logger) error { +func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, cfg *tls.Config, logger *zap.Logger, opts *config.Config) error { if err := conn.WritePacket(identity); err != nil { return fmt.Errorf("failed to send identity: %w", err) } @@ -468,7 +463,7 @@ func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *pr // shortly after a successful dial, e.g. a dying peer) keeps the // counter so the backoff keeps escalating instead of hammering at the // 2s floor forever. - if sender.ConnectionAge() >= reconnectFlapThreshold { + if sender.ConnectionAge() >= config.Duration(opts.Reconnect.FlapThreshold) { sender.ResetReconnectAttempt() } // Prevent multiple concurrent reconnect goroutines for the same device. @@ -477,7 +472,7 @@ func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *pr zap.String("device_id", sender.ID())) return } - go reconnectWithBackoff(ctx, sender, lastIP, identity, cfg, devices, plugins, localDeviceID, logger) + go reconnectWithBackoff(ctx, sender, lastIP, identity, cfg, devices, plugins, localDeviceID, logger, opts) } dev.Connect(ctx, conn, dispatch, onConnect, onDisconnect) @@ -502,8 +497,9 @@ func reconnectWithBackoff( plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, + opts *config.Config, ) { - const maxBackoff = 5 * time.Minute + maxBackoff := config.Duration(opts.Reconnect.MaxBackoff) attempt := dev.ReconnectAttempt() defer dev.ReconnectDone() @@ -534,7 +530,7 @@ func reconnectWithBackoff( return } - backoff := device.ReconnectBackoff(attempt, maxBackoff) + backoff := device.ReconnectBackoff(attempt, maxBackoff, config.Duration(opts.Reconnect.InitialBackoff)) logger.Debug("auto-reconnect: waiting before next attempt", zap.String("device_id", dev.ID()), zap.Int("attempt", attempt+1), @@ -561,11 +557,8 @@ func reconnectWithBackoff( // Prefer the peer's last advertised listening port over the // default: the identity may carry a non-standard port (or none // at all, in which case LastPort is 0 and we fall back). - port := 1716 - if p := dev.LastPort(); validDialPort(p) { - port = p - } - DialDevice(ctx, ip, port, dev.ID(), protocol.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false) + port := reconnectPort(dev, opts.TCPPort) + DialDevice(ctx, ip, port, dev.ID(), protocol.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) if dev.IsConnected() { logger.Info("auto-reconnect: succeeded", @@ -583,3 +576,11 @@ func reconnectWithBackoff( attempt++ } } + +// reconnectPort prefers the authenticated peer port, falling back to local configuration. +func reconnectPort(dev *device.Device, fallback int) int { + if p := dev.LastPort(); validDialPort(p) { + return p + } + return fallback +} diff --git a/internal/device/configurability_test.go b/internal/device/configurability_test.go new file mode 100644 index 0000000..b05fba3 --- /dev/null +++ b/internal/device/configurability_test.go @@ -0,0 +1,38 @@ +package device + +import ( + "go.uber.org/zap" + "testing" + "time" +) + +func TestConfiguredReconnectBackoff(t *testing.T) { + for _, tc := range []struct { + attempt int + initial, cap, want time.Duration + }{ + {0, time.Second, 5 * time.Second, time.Second}, + {2, time.Second, 5 * time.Second, 4 * time.Second}, + {3, time.Second, 5 * time.Second, 5 * time.Second}, + {1000000, time.Nanosecond, time.Duration(1<<63 - 1), time.Duration(1<<63 - 1)}, + {1, time.Duration(1 << 62), time.Duration(1<<63 - 1), time.Duration(1<<63 - 1)}, + } { + if got := ReconnectBackoff(tc.attempt, tc.cap, tc.initial); got != tc.want { + t.Errorf("%+v: got %v", tc, got) + } + } +} + +func TestConfiguredPairIntentTTL(t *testing.T) { + dev := NewDevice("peer", "Peer", "phone", zap.NewNop()) + before := time.Now() + dev.RequestPairDial(time.Hour) + deadline := time.Unix(0, dev.pairIntentUntil.Load()) + if deadline.Before(before.Add(time.Hour)) || deadline.After(time.Now().Add(time.Hour)) { + t.Fatalf("unexpected intent deadline %s", deadline) + } + dev.ClearPairDial() + if dev.PairDialActive() { + t.Fatal("cleared intent still active") + } +} diff --git a/internal/device/device.go b/internal/device/device.go index c352ead..fbdbc6d 100644 --- a/internal/device/device.go +++ b/internal/device/device.go @@ -456,9 +456,13 @@ const pairDialIntentTTL = 5 * time.Minute // discovery announcement and arms the keep-alive intent until pairing starts, // is rejected, succeeds, or the TTL expires. Used when the user explicitly // runs `kcd pair ` for a device with no active connection. -func (d *Device) RequestPairDial() { +func (d *Device) RequestPairDial(intentTTL ...time.Duration) { + ttl := pairDialIntentTTL + if len(intentTTL) > 0 && intentTTL[0] > 0 { + ttl = intentTTL[0] + } d.pairDialRequested.Store(true) - d.pairIntentUntil.Store(time.Now().Add(pairDialIntentTTL).UnixNano()) + d.pairIntentUntil.Store(time.Now().Add(ttl).UnixNano()) } // ConsumePairDial reports and clears a pending explicit pair-dial request. diff --git a/internal/device/registry.go b/internal/device/registry.go index 0f4dbe4..1992648 100644 --- a/internal/device/registry.go +++ b/internal/device/registry.go @@ -115,20 +115,20 @@ func (r *Registry) Prune(threshold time.Duration) int { // ReconnectBackoff calculates exponential backoff duration based on attempts. // Caps out at maxDuration. -func ReconnectBackoff(attempt int, maxDuration time.Duration) time.Duration { - if attempt <= 0 { - return 2 * time.Second +func ReconnectBackoff(attempt int, maxDuration time.Duration, initial ...time.Duration) time.Duration { + base := 2 * time.Second + if len(initial) > 0 && initial[0] > 0 { + base = initial[0] } - // Guard against shift overflow. On 64-bit, 1<= 54 (producing 0 due to modular multiplication), and - // on 32-bit, 1<<32 wraps to 0. Cap early — the backoff is already - // enormous at attempt=30 (~68 years if uncapped). - if attempt >= 31 { - return maxDuration + if maxDuration <= 0 { + return 0 } - dur := time.Duration(1< maxDuration || dur <= 0 { - return maxDuration + delay := min(base, maxDuration) + for i := 0; i < attempt && delay < maxDuration; i++ { + if delay > maxDuration/2 { + return maxDuration + } + delay *= 2 } - return dur + return delay } diff --git a/internal/discovery/discovery.go b/internal/discovery/discovery.go index 8b79fb8..67e7a03 100644 --- a/internal/discovery/discovery.go +++ b/internal/discovery/discovery.go @@ -27,6 +27,7 @@ const ( type BroadcasterController struct { identityPacket *protocol.Packet interval time.Duration + idleInterval time.Duration shouldReduce func() bool logger *zap.Logger @@ -37,10 +38,15 @@ type BroadcasterController struct { } // NewBroadcasterController creates a controller that starts in stopped state. -func NewBroadcasterController(identity *protocol.Packet, interval time.Duration, logger *zap.Logger, shouldReduce func() bool) *BroadcasterController { +func NewBroadcasterController(identity *protocol.Packet, interval time.Duration, logger *zap.Logger, shouldReduce func() bool, idleInterval ...time.Duration) *BroadcasterController { + idle := 60 * time.Second + if len(idleInterval) > 0 && idleInterval[0] > 0 { + idle = idleInterval[0] + } return &BroadcasterController{ identityPacket: identity, interval: interval, + idleInterval: idle, shouldReduce: shouldReduce, logger: logger.With(zap.String("component", "broadcaster")), owners: make(map[string]struct{}), @@ -76,6 +82,7 @@ func (bc *BroadcasterController) StartOwned(parentCtx context.Context, owner str b := &Broadcaster{ identityPacket: bc.identityPacket, interval: bc.interval, + idleInterval: bc.idleInterval, logger: bc.logger, } go func() { @@ -149,6 +156,7 @@ func AdvertiseMDNS(ctx context.Context, identityPacket *protocol.Packet, logger type Broadcaster struct { identityPacket *protocol.Packet interval time.Duration + idleInterval time.Duration logger *zap.Logger } @@ -167,7 +175,10 @@ func NewBroadcaster(identity *protocol.Packet, interval time.Duration, logger *z // (mDNS advertisement is no longer tied to this loop — see AdvertiseMDNS.) func (b *Broadcaster) Run(ctx context.Context, shouldReduce func() bool) { normalInterval := b.interval - reducedInterval := 60 * time.Second + reducedInterval := b.idleInterval + if reducedInterval <= 0 { + reducedInterval = 60 * time.Second + } conn, err := net.ListenUDP("udp4", nil) if err != nil { diff --git a/internal/discovery/discovery_test.go b/internal/discovery/discovery_test.go index 5011171..75eee92 100644 --- a/internal/discovery/discovery_test.go +++ b/internal/discovery/discovery_test.go @@ -18,6 +18,16 @@ func testIdentity(t *testing.T) *protocol.Packet { return pkt } +func TestConfiguredBroadcastIntervals(t *testing.T) { + bc := NewBroadcasterController(testIdentity(t), 7*time.Second, zap.NewNop(), nil, 19*time.Second) + if bc.interval != 7*time.Second || bc.idleInterval != 19*time.Second { + t.Fatal("configured intervals not stored") + } + if bc.IsRunning() { + t.Fatal("configuration must not start broadcasts") + } +} + // Pairing and reconnect needs share one loop but must not cancel each // other: withdrawing one owner leaves the loop up while the other holds it, // and the loop stops only when the last owner withdraws. diff --git a/internal/ipc/configurability_test.go b/internal/ipc/configurability_test.go new file mode 100644 index 0000000..07c2f8b --- /dev/null +++ b/internal/ipc/configurability_test.go @@ -0,0 +1,32 @@ +package ipc + +import ( + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/plugins/pair" + "go.uber.org/zap" + "strings" + "testing" + "time" +) + +func TestConfiguredPairListenTimeout(t *testing.T) { + logger := zap.NewNop() + bus := events.NewBus(logger) + devices := device.NewRegistry(bus) + cfg := config.Defaults() + pl := pair.NewPairPlugin(devices, nil, cfg.Pairing, nil, bus, logger) + h := NewHandler(devices, nil, pl, "", bus, 0) + h.SetPairListenTimeout(time.Millisecond) + done := make(chan Response, 1) + go func() { done <- h.handlePairListen() }() + select { + case resp := <-done: + if resp.OK || !strings.Contains(resp.Error, "(1ms)") { + t.Fatalf("unexpected response: %+v", resp) + } + case <-time.After(time.Second): + t.Fatal("configured timeout not applied") + } +} diff --git a/internal/ipc/handler.go b/internal/ipc/handler.go index 2e320bc..42674cf 100644 --- a/internal/ipc/handler.go +++ b/internal/ipc/handler.go @@ -2,6 +2,7 @@ package ipc import ( "encoding/json" + "fmt" "time" "github.com/bethropolis/kcd/internal/device" @@ -14,13 +15,14 @@ import ( // Handler handles incoming IPC requests. type Handler struct { - devices *device.Registry - plugins *plugin.Registry - pairPlugin *pair.PairPlugin - statePath string - bus *events.Bus - routes map[string]func(Request) Response - pruneThreshold time.Duration + devices *device.Registry + plugins *plugin.Registry + pairPlugin *pair.PairPlugin + statePath string + bus *events.Bus + routes map[string]func(Request) Response + pruneThreshold time.Duration + pairListenTimeout time.Duration // pairDialHook, when set, dials a disconnected device on explicit user // pair request (`kcd pair `). The daemon wires this to DialDevice // using the device's last-seen discovery address, so pairing an unpaired @@ -213,7 +215,18 @@ func (h *Handler) forgetContacts(deviceID string) { } } +// SetPairListenTimeout configures the wait window before the server starts. +func (h *Handler) SetPairListenTimeout(timeout time.Duration) { + if timeout > 0 { + h.pairListenTimeout = timeout + } +} + func (h *Handler) handlePairListen() Response { + timeout := h.pairListenTimeout + if timeout <= 0 { + timeout = 60 * time.Second + } // Report any device already in StatePairRequestedByPeer WITHOUT // accepting it. The caller (CLI / GUI / script) inspects the candidate // and decides: accept via CmdPair, reject via CmdUnpair. Auto-accepting @@ -248,8 +261,8 @@ func (h *Handler) handlePairListen() Response { } } return h.pairListenResult(dev, vKey) - case <-time.After(60 * time.Second): - return Response{OK: false, Error: "timed out waiting for pair request (60s)"} + case <-time.After(timeout): + return Response{OK: false, Error: fmt.Sprintf("timed out waiting for pair request (%s)", timeout)} } } diff --git a/internal/plugins/battery/battery.go b/internal/plugins/battery/battery.go index d3e787e..ac8568b 100644 --- a/internal/plugins/battery/battery.go +++ b/internal/plugins/battery/battery.go @@ -27,17 +27,23 @@ const ( // BatteryPlugin handles incoming battery state updates. type BatteryPlugin struct { - cfg config.BatteryConfig - bus *events.Bus - logger *zap.Logger + notifications config.NotificationConfig + cfg config.BatteryConfig + bus *events.Bus + logger *zap.Logger } // NewBatteryPlugin creates a BatteryPlugin. -func NewBatteryPlugin(cfg config.BatteryConfig, bus *events.Bus, logger *zap.Logger) *BatteryPlugin { +func NewBatteryPlugin(cfg config.BatteryConfig, bus *events.Bus, logger *zap.Logger, notifications ...config.NotificationConfig) *BatteryPlugin { + var notificationCfg config.NotificationConfig + if len(notifications) > 0 { + notificationCfg = notifications[0] + } return &BatteryPlugin{ - cfg: cfg, - bus: bus, - logger: logger.With(zap.String("plugin", "battery")), + notifications: notificationCfg, + cfg: cfg, + bus: bus, + logger: logger.With(zap.String("plugin", "battery")), } } @@ -115,7 +121,7 @@ func (p *BatteryPlugin) handleThreshold(dev device.Sender, body BatteryBody) { if message != "" { // Desktop notification. plugin.RunCommandAsync(p.logger, "notify-send", - "-a", "KDE Connect", + "-a", p.notifications.AppName(), "-u", urgency, "-i", "battery", dev.Name(), message, diff --git a/internal/plugins/clipboard/clipboard.go b/internal/plugins/clipboard/clipboard.go index e779e0b..d8a908a 100644 --- a/internal/plugins/clipboard/clipboard.go +++ b/internal/plugins/clipboard/clipboard.go @@ -19,6 +19,7 @@ import ( "github.com/bethropolis/kcd/internal/cert" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) @@ -32,6 +33,7 @@ const ( // ClipboardPlugin handles clipboard sync both directions. type ClipboardPlugin struct { + sidechannel transport.SidechannelOptions pushOnConnect bool lastTimestamp int64 tlsConfig *tls.Config @@ -45,11 +47,16 @@ type ClipboardPlugin struct { } // NewClipboardPlugin creates a clipboard plugin. -func NewClipboardPlugin(tlsConfig *tls.Config, logger *zap.Logger, pushOnConnect bool) *ClipboardPlugin { +func NewClipboardPlugin(tlsConfig *tls.Config, logger *zap.Logger, pushOnConnect bool, options ...transport.SidechannelOptions) *ClipboardPlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } if logger == nil { logger = zap.NewNop() } return &ClipboardPlugin{ + sidechannel: sidechannel, tlsConfig: tlsConfig, pushOnConnect: pushOnConnect, logger: logger.With(zap.String("plugin", "clipboard")), @@ -336,7 +343,7 @@ func (p *ClipboardPlugin) handleClipboardFile(ctx context.Context, dev device.Se tmpFile.Close() defer os.Remove(tmpPath) - if err := downloadToFile(ctx, remoteIP, payloadPort, payloadSize, tmpPath, p.tlsConfig, expectedFP, p.logger); err != nil { + if err := downloadToFile(ctx, remoteIP, payloadPort, payloadSize, tmpPath, p.tlsConfig, expectedFP, p.logger, p.sidechannel); err != nil { p.logger.Error("clipboard file: download failed", zap.Error(err)) return } @@ -371,30 +378,13 @@ func (p *ClipboardPlugin) handleClipboardFile(ctx context.Context, dev device.Se } // downloadToFile dials a TLS side-channel and streams the payload to dest. -func downloadToFile(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger) error { - addr := fmt.Sprintf("%s:%d", ip.String(), port) - dialer := &tls.Dialer{ - NetDialer: &net.Dialer{ - Timeout: 10 * time.Second, - KeepAlive: 30 * time.Second, - }, - Config: tlsConfig, - } - conn, err := dialer.DialContext(ctx, "tcp", addr) +func downloadToFile(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger, options ...transport.SidechannelOptions) error { + conn, err := transport.DialSidechannel(ctx, ip, port, tlsConfig, expectedFP, logger, options...) if err != nil { - return fmt.Errorf("clipboard: dial %s: %w", addr, err) + return err } defer conn.Close() - if tlsConn, ok := conn.(*tls.Conn); !ok { - return fmt.Errorf("clipboard: side-channel connection is not TLS") - } else if expectedFP == "" { - logger.Warn("clipboard: no pinned peer fingerprint, skipping side-channel verification", - zap.String("remote_addr", addr)) - } else if err := cert.VerifySideChannelPeer(tlsConn.ConnectionState(), expectedFP); err != nil { - return fmt.Errorf("clipboard: side-channel peer verification failed: %w", err) - } - f, err := os.Create(dest) if err != nil { return fmt.Errorf("clipboard: create file %s: %w", dest, err) diff --git a/internal/plugins/contacts/contacts.go b/internal/plugins/contacts/contacts.go index 6c7d5a7..5934ccd 100644 --- a/internal/plugins/contacts/contacts.go +++ b/internal/plugins/contacts/contacts.go @@ -87,13 +87,16 @@ type ContactsPlugin struct { // NewContactsPlugin creates a contacts plugin caching under // $XDG_DATA_HOME/kcd/contacts (0600 files, 0700 dirs). -func NewContactsPlugin(bus *events.Bus, logger *zap.Logger) *ContactsPlugin { +func NewContactsPlugin(bus *events.Bus, logger *zap.Logger, cacheDirs ...string) *ContactsPlugin { dataHome := os.Getenv("XDG_DATA_HOME") if dataHome == "" { home, _ := os.UserHomeDir() dataHome = filepath.Join(home, ".local", "share") } baseDir := filepath.Join(dataHome, "kcd", "contacts") + if len(cacheDirs) > 0 && cacheDirs[0] != "" { + baseDir = cacheDirs[0] + } _ = os.MkdirAll(baseDir, 0700) return &ContactsPlugin{ diff --git a/internal/plugins/mpris/artcache.go b/internal/plugins/mpris/artcache.go index e0d157d..a42cd85 100644 --- a/internal/plugins/mpris/artcache.go +++ b/internal/plugins/mpris/artcache.go @@ -31,12 +31,15 @@ type ArtCache struct { } // NewArtCache creates the cache directory and returns an empty cache. -func NewArtCache(logger *zap.Logger) *ArtCache { +func NewArtCache(logger *zap.Logger, cacheDirs ...string) *ArtCache { base, err := os.UserCacheDir() if err != nil || base == "" { base = filepath.Join(os.TempDir(), "kcd-cache") } dir := filepath.Join(base, "kcd", "art") + if len(cacheDirs) > 0 && cacheDirs[0] != "" { + dir = cacheDirs[0] + } if err := os.MkdirAll(dir, 0700); err != nil { logger.Warn("mpris: failed to create album art cache dir", zap.String("path", dir), zap.Error(err)) diff --git a/internal/plugins/mpris/mpris.go b/internal/plugins/mpris/mpris.go index 40438cb..c9d5b0e 100644 --- a/internal/plugins/mpris/mpris.go +++ b/internal/plugins/mpris/mpris.go @@ -63,7 +63,7 @@ type remotePositionTracker struct { playing bool } -func NewMPRISPlugin(tlsConfig *tls.Config, bus *events.Bus, pauseMusic bool, logger *zap.Logger) *MPRISPlugin { +func NewMPRISPlugin(tlsConfig *tls.Config, bus *events.Bus, pauseMusic bool, logger *zap.Logger, cacheDirs ...string) *MPRISPlugin { dbusConn, err := dbus.ConnectSessionBus() if err != nil { logger.Warn("mpris: failed to connect to D-Bus session bus", zap.Error(err)) @@ -86,7 +86,7 @@ func NewMPRISPlugin(tlsConfig *tls.Config, bus *events.Bus, pauseMusic bool, log remoteStateTimes: make(map[string]time.Time), positionTrackers: make(map[string]*remotePositionTracker), callPausedPlayers: make([]string, 0), - artCache: NewArtCache(logger), + artCache: NewArtCache(logger, cacheDirs...), } // Start the watcher immediately (like C++ does in constructor). diff --git a/internal/plugins/notification/notification.go b/internal/plugins/notification/notification.go index a510634..b768f1a 100644 --- a/internal/plugins/notification/notification.go +++ b/internal/plugins/notification/notification.go @@ -21,11 +21,13 @@ import ( "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) // NotificationPlugin handles incoming notifications and displays them on the desktop. type NotificationPlugin struct { + sidechannel transport.SidechannelOptions bus *events.Bus tlsConfig *tls.Config logger *zap.Logger @@ -42,13 +44,18 @@ type NotificationPlugin struct { // NewNotificationPlugin creates a NotificationPlugin. // tlsConfig is used to fetch notification icon payloads over the KDE Connect // side-channel; pass nil to disable icon fetching. -func NewNotificationPlugin(cfg config.NotificationPluginConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger) *NotificationPlugin { +func NewNotificationPlugin(cfg config.NotificationPluginConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...transport.SidechannelOptions) *NotificationPlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } p := &NotificationPlugin{ - cfg: cfg, - bus: bus, - tlsConfig: tlsConfig, - logger: logger.With(zap.String("plugin", "notification")), - newExec: exec.CommandContext, + sidechannel: sidechannel, + cfg: cfg, + bus: bus, + tlsConfig: tlsConfig, + logger: logger.With(zap.String("plugin", "notification")), + newExec: exec.CommandContext, } // Probe --print-id support by checking --help output. @@ -318,28 +325,13 @@ func (p *NotificationPlugin) fetchIcon( return "" } - addr := fmt.Sprintf("%s:%d", remoteIP, port) - dialer := &tls.Dialer{ - NetDialer: &net.Dialer{Timeout: 10 * time.Second}, - Config: p.tlsConfig, - } - conn, err := dialer.DialContext(ctx, "tcp", addr) + conn, err := transport.DialSidechannel(ctx, remoteIP, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) if err != nil { p.logger.Debug("notification: icon dial failed", zap.Error(err)) return "" } defer conn.Close() - if tlsConn, ok := conn.(*tls.Conn); !ok { - p.logger.Debug("notification: icon connection is not TLS") - return "" - } else if expectedFP == "" { - p.logger.Warn("notification: no pinned peer fingerprint, skipping icon verification") - } else if err := cert.VerifySideChannelPeer(tlsConn.ConnectionState(), expectedFP); err != nil { - p.logger.Warn("notification: icon peer verification failed, refusing download", zap.Error(err)) - return "" - } - f, err := os.Create(iconPath) if err != nil { return "" diff --git a/internal/plugins/ping/ping.go b/internal/plugins/ping/ping.go index 8a51438..99c6852 100644 --- a/internal/plugins/ping/ping.go +++ b/internal/plugins/ping/ping.go @@ -19,7 +19,14 @@ type PingPlugin struct { logger *zap.Logger } -func NewPingPlugin(cfg config.PingConfig, bus *events.Bus, logger *zap.Logger) *PingPlugin { +func NewPingPlugin(cfg config.PingConfig, bus *events.Bus, logger *zap.Logger, notifications ...config.NotificationConfig) *PingPlugin { + if cfg.AppName == "" { + var notificationCfg config.NotificationConfig + if len(notifications) > 0 { + notificationCfg = notifications[0] + } + cfg.AppName = notificationCfg.AppName() + } return &PingPlugin{ cfg: cfg, bus: bus, diff --git a/internal/plugins/share/share.go b/internal/plugins/share/share.go index 9b793a2..4c4da88 100644 --- a/internal/plugins/share/share.go +++ b/internal/plugins/share/share.go @@ -19,6 +19,7 @@ import ( "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/plugin" "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) @@ -66,6 +67,7 @@ func (t *progressThrottle) Update(current, _ int64) { } type SharePlugin struct { + sidechannel transport.SidechannelOptions DownloadDir string cfg config.ShareConfig TLSConfig *tls.Config @@ -73,8 +75,13 @@ type SharePlugin struct { bus *events.Bus } -func NewSharePlugin(downloadDir string, cfg config.ShareConfig, tlsConfig *tls.Config, bus *events.Bus, logger *zap.Logger) *SharePlugin { +func NewSharePlugin(downloadDir string, cfg config.ShareConfig, tlsConfig *tls.Config, bus *events.Bus, logger *zap.Logger, options ...transport.SidechannelOptions) *SharePlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } return &SharePlugin{ + sidechannel: sidechannel, DownloadDir: downloadDir, cfg: cfg, TLSConfig: tlsConfig, @@ -184,7 +191,7 @@ func (p *SharePlugin) Handle(ctx context.Context, dev device.Sender, pkt *protoc onProgress = throttle.Update } - err := ReceiveSideChannel(context.Background(), remoteIP, payloadPort, payloadSize, destPath, p.TLSConfig, expectedFP, onProgress, p.Logger) + err := ReceiveSideChannel(context.Background(), remoteIP, payloadPort, payloadSize, destPath, p.TLSConfig, expectedFP, onProgress, p.Logger, p.sidechannel) if err != nil { p.Logger.Error("share receive failed", zap.Error(err)) if p.bus != nil { diff --git a/internal/plugins/share/transfer.go b/internal/plugins/share/transfer.go index 9f35f75..e4bdb1a 100644 --- a/internal/plugins/share/transfer.go +++ b/internal/plugins/share/transfer.go @@ -12,6 +12,7 @@ import ( "github.com/bethropolis/kcd/internal/cert" "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) @@ -30,43 +31,17 @@ func (pw *progressWriter) Write(p []byte) (int, error) { return n, nil } -func ReceiveSideChannel(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, onProgress func(int64, int64), logger *zap.Logger) error { +func ReceiveSideChannel(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, onProgress func(int64, int64), logger *zap.Logger, options ...transport.SidechannelOptions) error { if size < 0 { return fmt.Errorf("share: indefinite payload sizes (-1) are not supported") } - addr := fmt.Sprintf("%s:%d", ip.String(), port) - - dialer := &tls.Dialer{ - NetDialer: &net.Dialer{ - Timeout: 10 * time.Second, - KeepAlive: 30 * time.Second, - }, - Config: tlsConfig, - } - conn, err := dialer.DialContext(ctx, "tcp", addr) + conn, err := transport.DialSidechannel(ctx, ip, port, tlsConfig, expectedFP, logger, options...) if err != nil { - return fmt.Errorf("share: dial side-channel %s: %w", addr, err) + return err } defer conn.Close() - // The phone's TLS cert isn't verified against the paired fingerprint by - // default (self-signed), so confirm the peer is the device we expect - // before pulling bytes from it. - if tlsConn, ok := conn.(*tls.Conn); !ok { - return fmt.Errorf("share: side-channel connection is not TLS") - } else if expectedFP == "" { - logger.Warn("share: no pinned peer fingerprint, skipping side-channel verification", - zap.String("remote_addr", addr)) - } else if err := cert.VerifySideChannelPeer(tlsConn.ConnectionState(), expectedFP); err != nil { - logger.Error("share: side-channel peer verification failed", - zap.String("remote_addr", addr), - zap.String("expected_fp", expectedFP), - zap.Error(err), - ) - return fmt.Errorf("share: side-channel peer verification failed: %w", err) - } - f, err := os.OpenFile(dest, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600) if err != nil { return fmt.Errorf("share: create file %s: %w", dest, err) diff --git a/internal/plugins/sms/sms.go b/internal/plugins/sms/sms.go index 818c7ad..3c98d7b 100644 --- a/internal/plugins/sms/sms.go +++ b/internal/plugins/sms/sms.go @@ -9,7 +9,6 @@ import ( "net" "os" "path/filepath" - "strconv" "strings" "time" "unicode" @@ -20,6 +19,7 @@ import ( "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/plugin" "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) @@ -42,23 +42,41 @@ const ( // SMSPlugin implements SMS sending, receiving, conversation browsing, and MMS // attachment handling for KDE Connect. type SMSPlugin struct { - cfg config.SMSConfig - bus *events.Bus - tlsConfig *tls.Config - logger *zap.Logger - cacheDir string + sidechannel transport.SidechannelOptions + notifications config.NotificationConfig + cfg config.SMSConfig + bus *events.Bus + tlsConfig *tls.Config + logger *zap.Logger + cacheDir string } -func NewSMSPlugin(cfg config.SMSConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger) *SMSPlugin { - cacheDir := filepath.Join(os.TempDir(), "kcd", "sms-attachments") +// Options customizes storage, network timeouts, and desktop notification identity. +type Options struct { + CacheDir string + Sidechannel transport.SidechannelOptions + Notifications config.NotificationConfig +} + +func NewSMSPlugin(cfg config.SMSConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...Options) *SMSPlugin { + var opts Options + if len(options) > 0 { + opts = options[0] + } + cacheDir := opts.CacheDir + if cacheDir == "" { + cacheDir = filepath.Join(os.TempDir(), "kcd", "sms-attachments") + } _ = os.MkdirAll(cacheDir, 0700) return &SMSPlugin{ - cfg: cfg, - bus: bus, - tlsConfig: tlsConfig, - logger: logger.With(zap.String("plugin", "sms")), - cacheDir: cacheDir, + sidechannel: opts.Sidechannel, + notifications: opts.Notifications, + cfg: cfg, + bus: bus, + tlsConfig: tlsConfig, + logger: logger.With(zap.String("plugin", "sms")), + cacheDir: cacheDir, } } @@ -185,7 +203,7 @@ func (p *SMSPlugin) handleMessages(_ context.Context, dev device.Sender, pkt *pr } title := fmt.Sprintf("SMS from %s", sender) plugin.RunCommandAsync(p.logger, "notify-send", - "-a", "KDE Connect", + "-a", p.notifications.AppName(), "-i", "dialog-information", title, msgText, @@ -257,30 +275,12 @@ func (p *SMSPlugin) receiveAttachment(ctx context.Context, ip net.IP, port int, if size <= 0 || size > maxSMSAttachmentBytes { return fmt.Errorf("sms: refusing attachment with invalid size %d (limit %d)", size, maxSMSAttachmentBytes) } - addr := net.JoinHostPort(ip.String(), strconv.Itoa(port)) - - dialer := &tls.Dialer{ - NetDialer: &net.Dialer{ - Timeout: 10 * time.Second, - KeepAlive: 30 * time.Second, - }, - Config: p.tlsConfig, - } - conn, err := dialer.DialContext(ctx, "tcp", addr) + conn, err := transport.DialSidechannel(ctx, ip, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) if err != nil { return fmt.Errorf("sms: connect to attachment side-channel: %w", err) } defer conn.Close() - if tlsConn, ok := conn.(*tls.Conn); !ok { - return fmt.Errorf("sms: attachment side-channel is not TLS") - } else if expectedFP == "" { - p.logger.Warn("sms: no pinned peer fingerprint, skipping side-channel verification", - zap.String("remote_addr", addr)) - } else if err := cert.VerifySideChannelPeer(tlsConn.ConnectionState(), expectedFP); err != nil { - return fmt.Errorf("sms: side-channel peer verification failed: %w", err) - } - f, err := os.Create(destPath) if err != nil { return fmt.Errorf("sms: create attachment file: %w", err) diff --git a/internal/plugins/telephony/telephony.go b/internal/plugins/telephony/telephony.go index 7fcad4f..4767756 100644 --- a/internal/plugins/telephony/telephony.go +++ b/internal/plugins/telephony/telephony.go @@ -5,6 +5,7 @@ import ( "encoding/json" "time" + "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/plugin" @@ -13,14 +14,20 @@ import ( ) type TelephonyPlugin struct { - bus *events.Bus - logger *zap.Logger + notifications config.NotificationConfig + bus *events.Bus + logger *zap.Logger } -func NewTelephonyPlugin(bus *events.Bus, logger *zap.Logger) *TelephonyPlugin { +func NewTelephonyPlugin(bus *events.Bus, logger *zap.Logger, notifications ...config.NotificationConfig) *TelephonyPlugin { + var notificationCfg config.NotificationConfig + if len(notifications) > 0 { + notificationCfg = notifications[0] + } return &TelephonyPlugin{ - bus: bus, - logger: logger.With(zap.String("plugin", "telephony")), + notifications: notificationCfg, + bus: bus, + logger: logger.With(zap.String("plugin", "telephony")), } } @@ -76,7 +83,7 @@ func (p *TelephonyPlugin) Handle(ctx context.Context, dev device.Sender, pkt *pr return } - plugin.RunCommandAsync(p.logger, "notify-send", "-a", "KDE Connect", "-u", urgency, title, message) + plugin.RunCommandAsync(p.logger, "notify-send", "-a", p.notifications.AppName(), "-u", urgency, title, message) }() return nil diff --git a/internal/transport/sidechannel.go b/internal/transport/sidechannel.go new file mode 100644 index 0000000..f8dd766 --- /dev/null +++ b/internal/transport/sidechannel.go @@ -0,0 +1,70 @@ +package transport + +import ( + "context" + "crypto/tls" + "fmt" + "net" + "strconv" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "go.uber.org/zap" +) + +// SidechannelOptions bounds the entire connection establishment phase — TCP +// dial plus TLS handshake plus pin verification — with a single overall +// timeout (default 15s). Payload streaming is never time-bounded. +type SidechannelOptions struct { + Timeout time.Duration +} + +// DialSidechannel connects as a TLS client and verifies the paired certificate +// before returning any payload bytes. The caller owns and must close the result. +// Once established, payload streaming carries no deadline, so large transfers +// are never truncated mid-stream by the helper. +func DialSidechannel(ctx context.Context, ip net.IP, port int, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger, options ...SidechannelOptions) (net.Conn, error) { + if ip == nil || port < 1 || port > 65535 { + return nil, fmt.Errorf("side-channel: invalid peer address") + } + if tlsConfig == nil { + return nil, fmt.Errorf("side-channel: TLS configuration is required") + } + opts := SidechannelOptions{} + if len(options) > 0 { + opts = options[0] + } + if opts.Timeout <= 0 { + opts.Timeout = 15 * time.Second + } + addr := net.JoinHostPort(ip.String(), strconv.Itoa(port)) + // One wall-clock budget covers everything before payload streaming: + // the TCP connect and the TLS handshake together. + setupCtx, cancel := context.WithTimeout(ctx, opts.Timeout) + defer cancel() + dialer := net.Dialer{Timeout: opts.Timeout, KeepAlive: 30 * time.Second} + raw, err := dialer.DialContext(setupCtx, "tcp", addr) + if err != nil { + return nil, fmt.Errorf("side-channel: dial %s: %w", addr, err) + } + conn := tls.Client(raw, tlsConfig) + if err := conn.HandshakeContext(setupCtx); err != nil { + conn.Close() + return nil, fmt.Errorf("side-channel: handshake %s: %w", addr, err) + } + if expectedFP == "" { + if logger != nil { + logger.Warn("side-channel: no pinned peer fingerprint, skipping verification", zap.String("remote_addr", addr)) + } + } else if err := cert.VerifySideChannelPeer(conn.ConnectionState(), expectedFP); err != nil { + conn.Close() + return nil, fmt.Errorf("side-channel: peer verification failed: %w", err) + } + // The caller's context governs connection establishment only. Once the + // handshake is done, cancellation has no effect on payload I/O — the + // pre-helper plugins dialed with the request context but streamed + // unbounded, and download goroutines may legitimately outlive the + // packet-handler context. Callers own the returned conn and close it + // when streaming finishes. + return conn, nil +} diff --git a/internal/transport/sidechannel_test.go b/internal/transport/sidechannel_test.go new file mode 100644 index 0000000..9b50d5f --- /dev/null +++ b/internal/transport/sidechannel_test.go @@ -0,0 +1,183 @@ +package transport + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "io" + "net" + "strconv" + "sync" + "testing" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "go.uber.org/zap" +) + +// startSidechannelServer accepts one TLS connection and streams N bytes to it. +// It reports the peer certificate so the test can pin the fingerprint. +func startSidechannelServer(t *testing.T, tlsConfig *tls.Config, payload []byte, delay time.Duration) (addr string, fp string, done <-chan error) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { ln.Close() }) + + tlsCert := tlsConfig.Certificates[0] + leaf, err := x509.ParseCertificate(tlsCert.Certificate[0]) + if err != nil { + t.Fatalf("parse cert: %v", err) + } + fp = cert.Fingerprint(leaf) + + doneCh := make(chan error, 1) + done = doneCh + go func() { + conn, err := ln.Accept() + if err != nil { + doneCh <- err + return + } + defer conn.Close() + tlsConn := tls.Server(conn, tlsConfig) + if err := tlsConn.HandshakeContext(context.Background()); err != nil { + doneCh <- err + return + } + // Simulate a slow sender: the payload takes longer than any sane + // per-chunk idle deadline would allow, but is a healthy stream. + chunk := 32 * 1024 + sent := 0 + for sent < len(payload) { + n := chunk + if sent+n > len(payload) { + n = len(payload) - sent + } + if _, err := tlsConn.Write(payload[sent : sent+n]); err != nil { + doneCh <- err + return + } + sent += n + if delay > 0 { + time.Sleep(delay) + } + } + doneCh <- nil + }() + return ln.Addr().String(), fp, done +} + +// TestDialSidechannelRoundTrip verifies the helper streams a payload larger +// than the setup timeout window: the sidechannel timeout bounds connection +// establishment only, never payload I/O. +func TestDialSidechannelRoundTrip(t *testing.T) { + payload := make([]byte, 256*1024) + if _, err := rand.Read(payload); err != nil { + t.Fatalf("rand: %v", err) + } + tlsCert, err := cert.GenerateSelfSigned("sidechannel_test_device") + if err != nil { + t.Fatalf("generate cert: %v", err) + } + tlsConfig := cert.TLSConfig(tlsCert) + + // A 256 KiB payload at 50 ms/chunk takes ~400 ms of streaming — long + // after a setup timeout of 15s would have been exceeded if it applied + // to reads. The transfer must still complete. + addr, fp, serverDone := startSidechannelServer(t, tlsConfig, payload, 50*time.Millisecond) + host, portStr, _ := net.SplitHostPort(addr) + port, _ := strconv.Atoi(portStr) + + conn, err := DialSidechannel(context.Background(), net.ParseIP(host), port, tlsConfig, fp, zap.NewNop(), SidechannelOptions{Timeout: 2 * time.Second}) + if err != nil { + t.Fatalf("DialSidechannel: %v", err) + } + defer conn.Close() + + got, err := io.ReadAll(conn) + if err != nil { + t.Fatalf("read payload: %v", err) + } + if !bytes.Equal(got, payload) { + t.Fatalf("payload mismatch: got %d bytes, want %d", len(got), len(payload)) + } + if err := <-serverDone; err != nil { + t.Fatalf("server: %v", err) + } +} + +// TestDialSidechannelPinMismatch verifies the helper rejects a peer whose +// certificate does not match the pinned fingerprint, before any payload flows. +func TestDialSidechannelPinMismatch(t *testing.T) { + tlsCert, err := cert.GenerateSelfSigned("sidechannel_evil_device") + if err != nil { + t.Fatalf("generate cert: %v", err) + } + tlsConfig := cert.TLSConfig(tlsCert) + + payload := []byte("secret") + addr, _, _ := startSidechannelServer(t, tlsConfig, payload, 0) + host, portStr, _ := net.SplitHostPort(addr) + port, _ := strconv.Atoi(portStr) + + wrongFP := "0000000000000000000000000000000000000000000000000000000000000000" + _, err = DialSidechannel(context.Background(), net.ParseIP(host), port, tlsConfig, wrongFP, zap.NewNop(), SidechannelOptions{Timeout: 2 * time.Second}) + if err == nil { + t.Fatal("expected pin verification failure, got nil error") + } + t.Logf("rejected as expected: %v", err) +} + +// TestDialSidechannelSetupTimeout verifies the timeout bounds connection +// establishment: a peer that accepts TCP but never completes the TLS +// handshake must surface an error within roughly the configured budget. +func TestDialSidechannelSetupTimeout(t *testing.T) { + // Plain TCP listener: accepts but never speaks TLS. + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer ln.Close() + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + conn, err := ln.Accept() + if err != nil { + return + } + // Hold the connection open without any TLS response. + buf := make([]byte, 1024) + for { + if _, err := conn.Read(buf); err != nil { + return + } + } + }() + + tlsCert, err := cert.GenerateSelfSigned("sidechannel_timeout_device") + if err != nil { + t.Fatalf("generate cert: %v", err) + } + tlsConfig := cert.TLSConfig(tlsCert) + + host, portStr, _ := net.SplitHostPort(ln.Addr().String()) + port, _ := strconv.Atoi(portStr) + + start := time.Now() + _, err = DialSidechannel(context.Background(), net.ParseIP(host), port, tlsConfig, "", zap.NewNop(), SidechannelOptions{Timeout: 500 * time.Millisecond}) + elapsed := time.Since(start) + + if err == nil { + t.Fatal("expected timeout error, got nil") + } + if elapsed > 3*time.Second { + t.Fatalf("setup timeout not enforced: took %v", elapsed) + } + t.Logf("timed out as expected after %v: %v", elapsed, err) +} diff --git a/packaging/kcd.example.toml b/packaging/kcd.example.toml index dfff97c..9b22289 100644 --- a/packaging/kcd.example.toml +++ b/packaging/kcd.example.toml @@ -78,12 +78,6 @@ battery = true # X11: requires xclip. clipboard = true -[clipboard] -# Push the local clipboard to a device every time it connects -# or reconnects. Off by default to avoid spurious pushes on -# daemon restart or phone reconnection. -# push_on_connect = false - # Forward phone notifications to the desktop via notify-send. # Requires libnotify / a running notification daemon. notification = true @@ -169,6 +163,37 @@ presenter = true # kcd volume mute — Mute/unmute a sink remotesystemvolume = true +# ─── Connection timing and storage ──────────────────────────────────────────── +# Durations use Go syntax (e.g. "750ms", "10s", "5m") and must be positive. +# Restart the daemon after changing these settings. + +[network] +# dial_timeout = "5s" # control-channel TCP dial +# handshake_timeout = "10s" # identity/TLS handshake +# sidechannel_timeout = "15s" # side-channel connection setup + +[reconnect] +# initial_backoff = "2s" +# max_backoff = "5m" # must be >= initial_backoff +# flap_threshold = "15s" # minimum lifetime for a stable connection + +[discovery] +# broadcast_interval = "30s" +# broadcast_idle_interval = "60s" # must be >= broadcast_interval +# These intervals apply only while on-demand UDP broadcast is running; +# they do not enable permanent broadcast or change mDNS advertisement. + +[cache] +# Empty values preserve existing storage paths; use absolute paths to override. +# sms_attachments_dir = "" # default: system temp directory/kcd/sms-attachments +# album_art_dir = "" # default: $XDG_CACHE_HOME/kcd/art +# contacts_dir = "" # default: $XDG_DATA_HOME/kcd/contacts +# Changing a path does not migrate existing files. + +[clipboard] +# Push the local clipboard on connection/reconnection. Off by default. +# push_on_connect = false + # ─── RunCommand: exposed commands ───────────────────────────────────────────── # Shell commands that the RunCommand plugin makes available to your phone. # Trigger them from the KDE Connect Android app → your device → Run command. @@ -194,6 +219,7 @@ suspend = "systemctl suspend" # The special key "*" sets the default for all unmatched apps. # # [notifications] +# app_name = "KDE Connect" # desktop notification branding, not a filter # "com.whatsapp" = "show" # "com.google.android.gm" = "silent" # "*" = "show" @@ -248,7 +274,8 @@ suspend = "systemctl suspend" # ─── Ping: connectivity check ──────────────────────────────────────────────── [ping] -# app_name = "KDE Connect" +# app_name = "" # inherit [notifications].app_name; non-empty overrides +# A previously explicit "KDE Connect" remains an override of the global name. # icon = "smartphone" # default_message = "" # used if the phone sends an empty ping message @@ -256,6 +283,8 @@ suspend = "systemctl suspend" [pairing] # timeout_secs = 30 # how long to wait for a pairing response +# intent_ttl = "5m" # lifetime of explicit `kcd pair ` intent +# listen_timeout = "60s" # maximum wait in `kcd pair` listen mode # ─── SMS: notification settings ────────────────────────────────────────────── diff --git a/pkg/client/client.go b/pkg/client/client.go index 9968929..733eebb 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -19,6 +19,9 @@ import ( type Client struct { SocketPath string Timeout time.Duration + // PairListenTimeout is the client's deadline for a pair_listen call. + // Set it above the daemon's pairing.listen_timeout; zero defaults to 70s. + PairListenTimeout time.Duration } // Call dialed the daemon, sends a request, and returns the response. @@ -99,12 +102,14 @@ func (c *Client) Devices() ([]device.DeviceInfo, error) { // PairListen enters listen mode: waits for an incoming pair request, auto-accepts // it, and returns the paired device info. Blocks up to 60 seconds. func (c *Client) PairListen() (*ipc.PairListenResult, error) { - // Use a longer timeout for the listen operation - savedTimeout := c.Timeout - c.Timeout = 70 * time.Second - defer func() { c.Timeout = savedTimeout }() + // Copy the client so a long listen never changes concurrent calls' deadlines. + listenClient := *c + listenClient.Timeout = c.PairListenTimeout + if listenClient.Timeout <= 0 { + listenClient.Timeout = 70 * time.Second + } - resp, err := c.Call(ipc.CmdPairListen, nil) + resp, err := listenClient.Call(ipc.CmdPairListen, nil) if err != nil { return nil, err } From 37486e81c15bca47778cfea618d76dbf9eb50453 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:42:14 +0300 Subject: [PATCH 02/12] test(transport): side-channel streaming must survive caller ctx cancel Regression for the AfterFunc variant considered mid-review: a Handle-scoped\ncontext dies when the packet handler returns, while download goroutines\ncontinue streaming. ctx now bounds establishment only; pin the behavior. --- internal/transport/sidechannel_test.go | 39 ++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/internal/transport/sidechannel_test.go b/internal/transport/sidechannel_test.go index 9b50d5f..517dec9 100644 --- a/internal/transport/sidechannel_test.go +++ b/internal/transport/sidechannel_test.go @@ -181,3 +181,42 @@ func TestDialSidechannelSetupTimeout(t *testing.T) { } t.Logf("timed out as expected after %v: %v", elapsed, err) } + +// TestDialSidechannelSurvivesContextCancel pins the establishment-only +// semantics: cancelling the caller's ctx after a successful dial must not +// abort an in-flight payload stream (Handle-scoped contexts die as soon as +// the packet handler returns, while download goroutines keep streaming). +func TestDialSidechannelSurvivesContextCancel(t *testing.T) { + tlsCert, err := cert.GenerateSelfSigned("sidechannel_cancel_device") + if err != nil { + t.Fatalf("generate cert: %v", err) + } + tlsConfig := cert.TLSConfig(tlsCert) + + payload := make([]byte, 128*1024) + if _, err := rand.Read(payload); err != nil { + t.Fatalf("rand: %v", err) + } + addr, fp, serverDone := startSidechannelServer(t, tlsConfig, payload, 20*time.Millisecond) + host, portStr, _ := net.SplitHostPort(addr) + port, _ := strconv.Atoi(portStr) + + ctx, cancel := context.WithCancel(context.Background()) + conn, err := DialSidechannel(ctx, net.ParseIP(host), port, tlsConfig, fp, zap.NewNop(), SidechannelOptions{Timeout: 2 * time.Second}) + if err != nil { + t.Fatalf("DialSidechannel: %v", err) + } + defer conn.Close() + cancel() // die like a Handle-scoped context would + + got, err := io.ReadAll(conn) + if err != nil { + t.Fatalf("stream must survive context cancellation: %v", err) + } + if !bytes.Equal(got, payload) { + t.Fatalf("payload mismatch: got %d bytes, want %d", len(got), len(payload)) + } + if err := <-serverDone; err != nil { + t.Fatalf("server: %v", err) + } +} From 23d27c8b37a783e87ee12af4334881e1ee7e9ca8 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:50:45 +0300 Subject: [PATCH 03/12] fix(pairing): match stock verification code algorithm Sort public-key DERs larger-first, append the initial request's ASCII timestamp for protocol v8, and display the first 8 hex chars uppercased, so the code shown by kcd pair matches the phone. Accept, reject and unpair packets no longer carry a fresh timestamp that would split the two sides' codes; the request generates and stores one timestamp and sends exactly what it stores. --- internal/cert/cert.go | 27 ++++++++++++++++------- internal/cert/cert_test.go | 39 ++++++++++++++++++++++++++++++++++ internal/ipc/handler.go | 4 ++-- internal/plugins/pair/pair.go | 36 +++++++++++++++++-------------- internal/protocol/pair.go | 11 +++++----- internal/protocol/pair_test.go | 33 ++++++++++++++++++++++++++++ 6 files changed, 119 insertions(+), 31 deletions(-) create mode 100644 internal/protocol/pair_test.go diff --git a/internal/cert/cert.go b/internal/cert/cert.go index ce98e5a..7edb31a 100644 --- a/internal/cert/cert.go +++ b/internal/cert/cert.go @@ -5,6 +5,7 @@ package cert import ( + "bytes" "crypto/rand" "crypto/rsa" "crypto/sha256" @@ -16,6 +17,7 @@ import ( "fmt" "math/big" "os" + "strconv" "strings" "time" ) @@ -156,9 +158,13 @@ func VerifySideChannelPeer(state tls.ConnectionState, expectedFP string) error { return nil } -// VerificationKey generates a verification fingerprint used for out-of-band pairing verification. -// It creates a SHA256 hash of the concatenated public keys to display a fingerprint. -func VerificationKey(localCert, remoteCert *x509.Certificate) string { +// VerificationKey generates the out-of-band pairing verification code +// displayed to the user. It matches the reference algorithm: the DER-encoded +// public keys are concatenated with the lexicographically larger one first, +// followed by the initial pair request's timestamp as an ASCII decimal +// string (protocol v8 and later; a non-positive timestamp hashes without it +// for older peers). The code is the first 8 hex characters, uppercased. +func VerificationKey(localCert, remoteCert *x509.Certificate, timestampSec int64) string { var localKey, remoteKey []byte if localCert != nil && localCert.PublicKey != nil { localKey, _ = x509.MarshalPKIXPublicKey(localCert.PublicKey) @@ -168,14 +174,19 @@ func VerificationKey(localCert, remoteCert *x509.Certificate) string { } var combined []byte - if string(localKey) < string(remoteKey) { - combined = append(localKey, remoteKey...) - } else { + if bytes.Compare(localKey, remoteKey) < 0 { combined = append(remoteKey, localKey...) + } else { + combined = append(localKey, remoteKey...) } - hash := sha256.Sum256(combined) - return hex.EncodeToString(hash[:]) + h := sha256.New() + h.Write(combined) + if timestampSec > 0 { + h.Write([]byte(strconv.FormatInt(timestampSec, 10))) + } + sum := h.Sum(nil) + return strings.ToUpper(hex.EncodeToString(sum[:4])) } // TLSConfig returns the standard tls.Config used for KDE Connect. diff --git a/internal/cert/cert_test.go b/internal/cert/cert_test.go index cf82f82..ca85728 100644 --- a/internal/cert/cert_test.go +++ b/internal/cert/cert_test.go @@ -3,6 +3,7 @@ package cert import ( "crypto/tls" "crypto/x509" + "encoding/hex" "os" "path/filepath" "testing" @@ -123,3 +124,41 @@ func TestVerifySideChannelPeer(t *testing.T) { t.Error("missing peer certificate accepted") } } + +func TestVerificationKeyMatchesReference(t *testing.T) { + // Fixed public-key DERs with an independently computed reference code + // (larger DER first + ASCII decimal timestamp, SHA-256, first 8 hex + // chars uppercased — the algorithm every stock client implements). + const ( + derA = "30820122300d06092a864886f70d01010105000382010f003082010a0282010100d0c0cb8bcf06d5f3ca9ba5d529d112764d79a7d0eeef2fb7522e027982b10dc3d5253961153e6a707032c3272818fe4e699f2d4d9289806a81f7613eeba87092a39206d756ed5317f47c565314b6b24a62d9d009fb8d9af1fccd1f387bb9f6cb8a6840cad05791310e25df77042121f8969e63b60381d722fa01d6cfb27a0e9b8716010b0da6037dacf5e7eb9f3b4f366d4dee37d976c76975f68751ae4677ed3aac6b3b5c885af7a65740c604abc17daf4fb260a35fcac7d9eff525a8febafb6c4239f6bc2f5655caeb1b979dd363a513b51d698b6b93c5396bfc6a9968aac07733d895e94da23080759822af22e54f778cb0f237ba6f12148d54d3d73819cf0203010001" + derB = "30820122300d06092a864886f70d01010105000382010f003082010a0282010100f7a42d36379fb98d5032b639a4f4d9d3360797831c70de6317f96f5c6bcc3325dd129d5816f7bbdff608ba71ae5c2a2de1516075d1129e34a6a1fa7558aeecdeceb5202e48863fa3453fa508891f5d21ca75d42120b82579102ccb250e55a7184bd8c50848aef917022b42a3d95403155c66d8b29996fbdd37444733347cbc33138112b6d7d993a965b955026a90f9caaaeacb5e519bcd2cecf38b012962af5bbe67b6c5b5ff203364f58ce39da6acfce86a07228c18d0280944677388f211e68a10938afeb0204d4545b031800373a86a9a7c5989ecf81fb5b59f3a6d68a7a781882f488fd0dab498383b86ce7fd85fbdc79a7bb71b5f606b6b6b86bea1fd910203010001" + timestamp = 1711234567 + want = "698179EF" + wantNoTS = "D7B8B287" + ) + mustPub := func(hexDER string) *x509.Certificate { + t.Helper() + raw, err := hex.DecodeString(hexDER) + if err != nil { + t.Fatalf("decode DER: %v", err) + } + pub, err := x509.ParsePKIXPublicKey(raw) + if err != nil { + t.Fatalf("parse public key: %v", err) + } + return &x509.Certificate{PublicKey: pub} + } + certA, certB := mustPub(derA), mustPub(derB) + + // Order-independent: both sides display the same code. + if got := VerificationKey(certA, certB, timestamp); got != want { + t.Errorf("VerificationKey(a, b) = %q, want %q", got, want) + } + if got := VerificationKey(certB, certA, timestamp); got != want { + t.Errorf("VerificationKey(b, a) = %q, want %q", got, want) + } + // Non-positive timestamp hashes without it (pre-v8 peers). + if got := VerificationKey(certA, certB, 0); got != wantNoTS { + t.Errorf("VerificationKey no-timestamp = %q, want %q", got, wantNoTS) + } +} diff --git a/internal/ipc/handler.go b/internal/ipc/handler.go index 42674cf..78fc722 100644 --- a/internal/ipc/handler.go +++ b/internal/ipc/handler.go @@ -157,7 +157,7 @@ func (h *Handler) handlePair(payload []byte) Response { } // Fallback if no pair plugin (shouldn't happen) - pkt, _ := protocol.NewPairPacket(protocol.PairAccept) + pkt, _ := protocol.NewPairPacket(protocol.PairAccept, 0) if err := dev.Send(pkt); err != nil { return Response{OK: false, Error: "failed to send pair packet"} } @@ -190,7 +190,7 @@ func (h *Handler) handleUnpair(payload []byte) Response { } // Fallback - pkt, _ := protocol.NewPairPacket(protocol.PairReject) + pkt, _ := protocol.NewPairPacket(protocol.PairReject, 0) _ = dev.Send(pkt) dev.Disconnect() h.devices.Remove(p.DeviceID) diff --git a/internal/plugins/pair/pair.go b/internal/plugins/pair/pair.go index dee2cff..b2c034d 100644 --- a/internal/plugins/pair/pair.go +++ b/internal/plugins/pair/pair.go @@ -46,6 +46,14 @@ func NewPairPlugin(devices *device.Registry, localCert *x509.Certificate, cfg co } } +// pairingTimestampFor returns the stored pair-request timestamp for a +// device, or zero if none was recorded (pre-v8 peer or unknown). +func (p *PairPlugin) pairingTimestampFor(deviceID string) int64 { + p.mu.Lock() + defer p.mu.Unlock() + return p.pairingTimestamp[deviceID] +} + // emit publishes an event to the bus if one is configured. func (p *PairPlugin) emit(typ events.EventType, dev *device.Device, vKey string) { if p.bus == nil { @@ -113,7 +121,7 @@ func (p *PairPlugin) handlePairRequest(_ context.Context, dev *device.Device, bo // The peer sends pair:true as confirmation/keep-alive. // Just acknowledge by sending pair:true back. p.logger.Debug("received pair confirmation from already paired device", zap.String("device_id", dev.ID())) - pkt, _ := protocol.NewPairPacket(protocol.PairAccept) + pkt, _ := protocol.NewPairPacket(protocol.PairAccept, 0) dev.Send(pkt) return nil @@ -129,7 +137,7 @@ func (p *PairPlugin) handlePairRequest(_ context.Context, dev *device.Device, bo zap.Int64("timestamp", body.Timestamp), zap.Int64("now", now)) // Send rejection - pkt, _ := protocol.NewPairPacket(protocol.PairReject) + pkt, _ := protocol.NewPairPacket(protocol.PairReject, 0) dev.Send(pkt) return nil } @@ -144,10 +152,7 @@ func (p *PairPlugin) handlePairRequest(_ context.Context, dev *device.Device, bo var vKey string peerCert := dev.PeerCert() if peerCert != nil { - vKey = cert.VerificationKey(p.localCert, peerCert) - if len(vKey) > 16 { - vKey = vKey[:16] - } + vKey = cert.VerificationKey(p.localCert, peerCert, p.pairingTimestampFor(dev.ID())) p.logger.Info("pairing verification code", zap.String("device_id", dev.ID()), zap.String("code", vKey)) @@ -211,7 +216,7 @@ func (p *PairPlugin) handleUnpairRequest(_ context.Context, dev *device.Device) // AcceptPairing accepts an incoming pair request. func (p *PairPlugin) AcceptPairing(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairAccept) + pkt, err := protocol.NewPairPacket(protocol.PairAccept, 0) if err != nil { return err } @@ -240,14 +245,16 @@ func (p *PairPlugin) RequestPairing(dev *device.Device) error { return p.AcceptPairing(dev) } - pkt, err := protocol.NewPairPacket(protocol.PairAccept) + // The request timestamp seeds the verification code on both sides, + // so generate it once and send exactly what we store. + timestamp := time.Now().Unix() + pkt, err := protocol.NewPairPacket(protocol.PairAccept, timestamp) if err != nil { return err } - // Store our timestamp p.mu.Lock() - p.pairingTimestamp[dev.ID()] = time.Now().Unix() + p.pairingTimestamp[dev.ID()] = timestamp p.mu.Unlock() if err := dev.Send(pkt); err != nil { @@ -257,10 +264,7 @@ func (p *PairPlugin) RequestPairing(dev *device.Device) error { peerCert := dev.PeerCert() if peerCert != nil { - vKey := cert.VerificationKey(p.localCert, peerCert) - if len(vKey) > 16 { - vKey = vKey[:16] - } + vKey := cert.VerificationKey(p.localCert, peerCert, timestamp) p.logger.Info("pairing verification code", zap.String("device_id", dev.ID()), zap.String("code", vKey)) @@ -278,7 +282,7 @@ func (p *PairPlugin) RequestPairing(dev *device.Device) error { // RejectPairing rejects an incoming pair request. func (p *PairPlugin) RejectPairing(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairReject) + pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) if err != nil { return err } @@ -302,7 +306,7 @@ func (p *PairPlugin) RejectPairing(dev *device.Device) error { // Unpair removes pairing with a device. func (p *PairPlugin) Unpair(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairReject) + pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) if err != nil { return err } diff --git a/internal/protocol/pair.go b/internal/protocol/pair.go index 097d828..ca0b4d7 100644 --- a/internal/protocol/pair.go +++ b/internal/protocol/pair.go @@ -1,7 +1,5 @@ package protocol -import "time" - // TypePair is the packet type for pairing requests/responses. const TypePair = "kdeconnect.pair" @@ -17,10 +15,13 @@ type PairBody struct { Timestamp int64 `json:"timestamp,omitempty"` } -// NewPairPacket creates a pairing packet (accept or reject). -func NewPairPacket(pair bool) (*Packet, error) { +// NewPairPacket creates a pairing packet. Only the initial pair request +// carries a timestamp; accept, reject and unpair packets omit it. The +// verification code on both sides derives from the request's timestamp, so +// sending a fresh one on accept would desynchronize the displayed codes. +func NewPairPacket(pair bool, timestamp int64) (*Packet, error) { return NewPacket(TypePair, PairBody{ Pair: pair, - Timestamp: time.Now().Unix(), + Timestamp: timestamp, }) } diff --git a/internal/protocol/pair_test.go b/internal/protocol/pair_test.go new file mode 100644 index 0000000..3a1464b --- /dev/null +++ b/internal/protocol/pair_test.go @@ -0,0 +1,33 @@ +package protocol + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestNewPairPacketTimestampRules(t *testing.T) { + // Accept, reject and unpair packets carry no timestamp; only the + // initial request does. Both sides derive the verification code from + // the request's timestamp, so a fresh stamp on accept would split them. + accept, err := NewPairPacket(PairAccept, 0) + if err != nil { + t.Fatalf("NewPairPacket accept: %v", err) + } + if strings.Contains(string(accept.Body), "timestamp") { + t.Errorf("accept packet must omit timestamp, got %s", accept.Body) + } + + const ts = int64(1711234567) + req, err := NewPairPacket(PairAccept, ts) + if err != nil { + t.Fatalf("NewPairPacket request: %v", err) + } + var body PairBody + if err := json.Unmarshal(req.Body, &body); err != nil { + t.Fatalf("unmarshal request body: %v", err) + } + if !body.Pair || body.Timestamp != ts { + t.Errorf("request body = %+v, want pair:true timestamp:%d", body, ts) + } +} From aaab821cdfbeeb68a04fd9d14804dd30af87b3c8 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:54:15 +0300 Subject: [PATCH 04/12] fix(sms): send v2 schema and accept numeric read flags MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Stock message apps read only messageBody with the recipient list in addresses and require version 2; the previous flat phoneNumber payload made the phone transmit a blank text. Inbound message batches carry the SQLite read column as 0/1, which strict bool decoding rejected and dropped whole threads — a shared tolerant boolean now covers both numeric and string spellings. --- internal/plugins/sms/sms.go | 30 ++++++++------ internal/plugins/sms/sms_test.go | 65 ++++++++++++++++++++++++++++++ internal/protocol/flexbool.go | 33 +++++++++++++++ internal/protocol/flexbool_test.go | 32 +++++++++++++++ 4 files changed, 147 insertions(+), 13 deletions(-) create mode 100644 internal/protocol/flexbool.go create mode 100644 internal/protocol/flexbool_test.go diff --git a/internal/plugins/sms/sms.go b/internal/plugins/sms/sms.go index 3c98d7b..8b8b86b 100644 --- a/internal/plugins/sms/sms.go +++ b/internal/plugins/sms/sms.go @@ -104,16 +104,16 @@ type SMSMessagesPacket struct { } type SMSMessage struct { - Event int `json:"event"` - Body string `json:"body"` - Addresses []SMSAddress `json:"addresses"` - Date int64 `json:"date"` - Type int `json:"type"` - ThreadID int64 `json:"thread_id"` - Read bool `json:"read"` - UID int64 `json:"u_id,omitempty"` - SubID int `json:"sub_id,omitempty"` - Attachments []SMSAttachment `json:"attachments,omitempty"` + Event int `json:"event"` + Body string `json:"body"` + Addresses []SMSAddress `json:"addresses"` + Date int64 `json:"date"` + Type int `json:"type"` + ThreadID int64 `json:"thread_id"` + Read protocol.FlexBool `json:"read"` + UID int64 `json:"u_id,omitempty"` + SubID int `json:"sub_id,omitempty"` + Attachments []SMSAttachment `json:"attachments,omitempty"` } type SMSAddress struct { @@ -185,7 +185,7 @@ func (p *SMSPlugin) handleMessages(_ context.Context, dev device.Sender, pkt *pr "date": msg.Date, "type": msg.Type, "thread_id": msg.ThreadID, - "read": msg.Read, + "read": bool(msg.Read), "event": msg.Event, "u_id": msg.UID, "sub_id": msg.SubID, @@ -298,10 +298,14 @@ func (p *SMSPlugin) receiveAttachment(ctx context.Context, ip net.IP, port int, // --- SMS sending ----------------------------------------------------------- func (p *SMSPlugin) SendSMS(dev device.Sender, phoneNumber, message string) error { + // v2 schema: the phone reads only messageBody, with addresses as the + // primary recipient list (phoneNumber stays as a legacy fallback for + // older peers). Without addresses/version the phone sends a blank SMS. body := map[string]any{ - "sendSms": true, - "phoneNumber": phoneNumber, + "version": 2, + "addresses": []map[string]string{{"address": phoneNumber}}, "messageBody": message, + "phoneNumber": phoneNumber, } pkt, err := protocol.NewPacket(PacketTypeSMSRequest, body) if err != nil { diff --git a/internal/plugins/sms/sms_test.go b/internal/plugins/sms/sms_test.go index 796bfa4..c7dbe19 100644 --- a/internal/plugins/sms/sms_test.go +++ b/internal/plugins/sms/sms_test.go @@ -2,10 +2,15 @@ package sms import ( "context" + "crypto/x509" + "encoding/json" + "net" "strings" "testing" "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" "go.uber.org/zap" ) @@ -41,3 +46,63 @@ func TestReceiveAttachmentRejectsBadSize(t *testing.T) { } } } + +// captureSender implements device.Sender and records outbound packets. +type captureSender struct { + sent []*protocol.Packet +} + +func (s *captureSender) ID() string { return "test-device" } +func (s *captureSender) Name() string { return "test" } +func (s *captureSender) SetName(string) {} +func (s *captureSender) State() device.PairingState { return device.StatePaired } +func (s *captureSender) SetState(device.PairingState) {} +func (s *captureSender) Send(p *protocol.Packet) error { s.sent = append(s.sent, p); return nil } +func (s *captureSender) IsConnected() bool { return true } +func (s *captureSender) RemoteIP() net.IP { return nil } +func (s *captureSender) PeerCert() *x509.Certificate { return nil } +func (s *captureSender) HasCapability(string) bool { return true } +func (s *captureSender) UpdateBattery(int, bool) {} +func (s *captureSender) GetBattery() (int, bool) { return 0, false } + +func TestSendSMSUsesV2Schema(t *testing.T) { + p := NewSMSPlugin(config.SMSConfig{}, nil, nil, zap.NewNop()) + dev := &captureSender{} + if err := p.SendSMS(dev, "+1234567890", "hello"); err != nil { + t.Fatalf("SendSMS: %v", err) + } + if len(dev.sent) != 1 { + t.Fatalf("sent %d packets, want 1", len(dev.sent)) + } + var body struct { + Version int `json:"version"` + Addresses []SMSAddress `json:"addresses"` + MessageBody string `json:"messageBody"` + } + if err := json.Unmarshal(dev.sent[0].Body, &body); err != nil { + t.Fatalf("unmarshal send body: %v", err) + } + if body.Version != 2 { + t.Errorf("version = %d, want 2", body.Version) + } + if len(body.Addresses) != 1 || body.Addresses[0].Address != "+1234567890" { + t.Errorf("addresses = %+v, want recipient", body.Addresses) + } + if body.MessageBody != "hello" { + t.Errorf("messageBody = %q, want %q", body.MessageBody, "hello") + } +} + +func TestMessagesBatchAcceptsIntReadFlag(t *testing.T) { + // Android serializes the SQLite read column as a raw number (0/1). + for _, read := range []string{"1", "0", "true", "false"} { + pkt := &protocol.Packet{ + Type: PacketTypeSMSMessages, + Body: json.RawMessage(`{"version":2,"messages":[{"event":1,"body":"hi","addresses":[{"address":"+1"}],"date":1711234567,"type":1,"thread_id":7,"read":` + read + `}]}`), + } + p := NewSMSPlugin(config.SMSConfig{}, nil, nil, zap.NewNop()) + if err := p.Handle(context.Background(), &captureSender{}, pkt); err != nil { + t.Errorf("Handle with read=%s: %v, want nil", read, err) + } + } +} diff --git a/internal/protocol/flexbool.go b/internal/protocol/flexbool.go new file mode 100644 index 0000000..3fc70e4 --- /dev/null +++ b/internal/protocol/flexbool.go @@ -0,0 +1,33 @@ +package protocol + +import ( + "encoding/json" + "fmt" + "strings" +) + +// FlexBool is a bool that tolerates the loose encodings stock clients +// emit: JSON booleans, 0/1 numbers, and "true"/"false"/"1"/"0" strings +// (Android sends SMS read flags as 0/1 and telephony cancel flags as +// the string "true"). It marshals back as a plain JSON boolean. +type FlexBool bool + +// UnmarshalJSON accepts booleans, numbers, and strings. +func (b *FlexBool) UnmarshalJSON(data []byte) error { + var v any + if err := json.Unmarshal(data, &v); err != nil { + return err + } + switch t := v.(type) { + case bool: + *b = FlexBool(t) + case float64: + *b = t != 0 + case string: + s := strings.ToLower(strings.TrimSpace(t)) + *b = s == "true" || s == "1" + default: + return fmt.Errorf("protocol: cannot unmarshal %T into FlexBool", v) + } + return nil +} diff --git a/internal/protocol/flexbool_test.go b/internal/protocol/flexbool_test.go new file mode 100644 index 0000000..3bc7c40 --- /dev/null +++ b/internal/protocol/flexbool_test.go @@ -0,0 +1,32 @@ +package protocol + +import ( + "encoding/json" + "testing" +) + +func TestFlexBoolUnmarshal(t *testing.T) { + cases := []struct { + in string + want bool + }{ + {`true`, true}, + {`false`, false}, + {`1`, true}, + {`0`, false}, + {`"true"`, true}, + {`"false"`, false}, + {`"1"`, true}, + {`"0"`, false}, + } + for _, tc := range cases { + var b FlexBool + if err := json.Unmarshal([]byte(tc.in), &b); err != nil { + t.Errorf("unmarshal %s: %v", tc.in, err) + continue + } + if bool(b) != tc.want { + t.Errorf("unmarshal %s = %v, want %v", tc.in, b, tc.want) + } + } +} From 8a09cf189b5219f4767350baf6aceaa68ccee360 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:55:23 +0300 Subject: [PATCH 05/12] fix(battery): answer state requests instead of storing them Stock peers request our battery with request:true inside a regular battery packet; the empty body decoded as a 0% not-charging update and wiped the real charge. Requests are now answered with the local state and never reach UpdateBattery, and the connect-time probe uses the packet shape stock peers honor. --- internal/plugins/battery/battery.go | 49 ++++++++++++++++-------- internal/plugins/battery/battery_test.go | 23 +++++++++++ 2 files changed, 56 insertions(+), 16 deletions(-) diff --git a/internal/plugins/battery/battery.go b/internal/plugins/battery/battery.go index ac8568b..dcebd9f 100644 --- a/internal/plugins/battery/battery.go +++ b/internal/plugins/battery/battery.go @@ -48,10 +48,13 @@ func NewBatteryPlugin(cfg config.BatteryConfig, bus *events.Bus, logger *zap.Log } // BatteryBody represents the body of a kdeconnect.battery packet. +// A body carrying only Request asks the peer to report its state; it is +// not a state update and must never touch the stored charge. type BatteryBody struct { CurrentCharge int `json:"currentCharge"` IsCharging bool `json:"isCharging"` ThresholdEvent int `json:"thresholdEvent"` + Request bool `json:"request,omitempty"` } func (p *BatteryPlugin) Name() string { return "Battery" } @@ -72,6 +75,13 @@ func (p *BatteryPlugin) Handle(ctx context.Context, dev device.Sender, pkt *prot return fmt.Errorf("battery: decode body: %w", err) } + if body.Request { + // Stock peers ask for our state with request:true inside a + // regular battery packet. Answer it; an empty body would + // otherwise clobber the stored charge to a bogus 0%. + return p.sendLocalState(dev) + } + dev.UpdateBattery(body.CurrentCharge, body.IsCharging) if body.ThresholdEvent != thresholdNone { @@ -79,25 +89,30 @@ func (p *BatteryPlugin) Handle(ctx context.Context, dev device.Sender, pkt *prot } case "kdeconnect.battery.request": - // The phone is asking for our local battery state. - charge, charging, err := readLocalBattery() - if err != nil { - p.logger.Debug("local battery unavailable, skipping response", zap.Error(err)) - return nil - } - pkt, err := protocol.NewPacket("kdeconnect.battery", BatteryBody{ - CurrentCharge: charge, - IsCharging: charging, - }) - if err != nil { - return fmt.Errorf("battery: create response packet: %w", err) - } - return dev.Send(pkt) + // Legacy request type kept for older peers. + return p.sendLocalState(dev) } return nil } +// sendLocalState reports the local battery to the peer. +func (p *BatteryPlugin) sendLocalState(dev device.Sender) error { + charge, charging, err := readLocalBattery() + if err != nil { + p.logger.Debug("local battery unavailable, skipping response", zap.Error(err)) + return nil + } + pkt, err := protocol.NewPacket("kdeconnect.battery", BatteryBody{ + CurrentCharge: charge, + IsCharging: charging, + }) + if err != nil { + return fmt.Errorf("battery: create response packet: %w", err) + } + return dev.Send(pkt) +} + func (p *BatteryPlugin) handleThreshold(dev device.Sender, body BatteryBody) { var message, urgency string @@ -140,8 +155,10 @@ func (p *BatteryPlugin) handleThreshold(dev device.Sender, body BatteryBody) { // OnConnect requests the phone's battery and sends our local battery state. func (p *BatteryPlugin) OnConnect(dev device.Sender) { - // Ask phone for its battery. - pkt, _ := protocol.NewPacket("kdeconnect.battery.request", map[string]any{ + // Ask phone for its battery. Stock peers only honor request:true + // inside a regular battery packet; the legacy request type is kept + // for backward compatibility on the inbound path. + pkt, _ := protocol.NewPacket("kdeconnect.battery", map[string]any{ "request": true, }) dev.Send(pkt) diff --git a/internal/plugins/battery/battery_test.go b/internal/plugins/battery/battery_test.go index 6b62a71..b1a1282 100644 --- a/internal/plugins/battery/battery_test.go +++ b/internal/plugins/battery/battery_test.go @@ -153,3 +153,26 @@ func TestBatteryPlugin_Handle_NoThreshold_NoEvent(t *testing.T) { // correct — no event } } + +func TestBatteryPlugin_Handle_RequestDoesNotClobberCharge(t *testing.T) { + logger := zaptest.NewLogger(t) + p, _ := newPlugin(t) + dev := device.NewDevice("dev1", "Test Phone", "phone", logger) + dev.UpdateBattery(85, true) + + // Stock peers ask for our state with request:true inside a regular + // battery packet. It must be answered, never stored as a 0% update. + pkt, _ := protocol.NewPacket("kdeconnect.battery", map[string]any{"request": true}) + if err := p.Handle(context.Background(), dev, pkt); err != nil { + t.Fatalf("Handle returned error: %v", err) + } + if charge, charging := dev.GetBattery(); charge != 85 || !charging { + t.Errorf("request packet clobbered charge: got (%d, %v), want (85, true)", charge, charging) + } + + // The connect-time exchange must not touch stored state either. + p.OnConnect(dev) + if charge, charging := dev.GetBattery(); charge != 85 || !charging { + t.Errorf("OnConnect clobbered charge: got (%d, %v), want (85, true)", charge, charging) + } +} From 8e50bb6efe75e672d4e67419c48bf1346ce89b47 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:56:11 +0300 Subject: [PATCH 06/12] fix(telephony): tolerate stock cancel and missed-call encodings Stock phones end calls with isCancel as the string "true", which strict bool decoding rejected and left the ringing state stuck, and report missed calls as missedCall, which matched nothing. Both spellings are now accepted and missedCall feeds the standard missed event so existing watchers keep working. --- internal/plugins/telephony/telephony.go | 21 +++-- internal/plugins/telephony/telephony_test.go | 80 ++++++++++++++++++++ 2 files changed, 93 insertions(+), 8 deletions(-) create mode 100644 internal/plugins/telephony/telephony_test.go diff --git a/internal/plugins/telephony/telephony.go b/internal/plugins/telephony/telephony.go index 4767756..5a6fa93 100644 --- a/internal/plugins/telephony/telephony.go +++ b/internal/plugins/telephony/telephony.go @@ -32,10 +32,10 @@ func NewTelephonyPlugin(bus *events.Bus, logger *zap.Logger, notifications ...co } type TelephonyBody struct { - Event string `json:"event"` // "ringing", "talking", "missed" - ContactName string `json:"contactName"` - PhoneNumber string `json:"phoneNumber"` - IsCancel bool `json:"isCancel"` + Event string `json:"event"` // "ringing", "talking", "missedCall" + ContactName string `json:"contactName"` + PhoneNumber string `json:"phoneNumber"` + IsCancel protocol.FlexBool `json:"isCancel"` } func (p *TelephonyPlugin) Name() string { return "Telephony" } @@ -51,16 +51,21 @@ func (p *TelephonyPlugin) Handle(ctx context.Context, dev device.Sender, pkt *pr return err } + canceled := bool(body.IsCancel) if p.bus != nil { - if body.IsCancel { + if canceled { p.bus.Publish(events.TypeTelephonyCanceled, dev.ID(), body) } else { - p.bus.Publish(events.EventType("telephony."+body.Event), dev.ID(), body) + eventType := events.EventType("telephony." + body.Event) + if body.Event == "missedCall" { + eventType = events.TypeTelephonyMissed + } + p.bus.Publish(eventType, dev.ID(), body) } } go func() { - if body.IsCancel { + if canceled { return } @@ -76,7 +81,7 @@ func (p *TelephonyPlugin) Handle(ctx context.Context, dev device.Sender, pkt *pr title = "📞 Incoming Call" message = "Ringing: " + caller urgency = "critical" - case "missed": + case "missed", "missedCall": title = "❌ Missed Call" message = "Missed call from " + caller default: diff --git a/internal/plugins/telephony/telephony_test.go b/internal/plugins/telephony/telephony_test.go new file mode 100644 index 0000000..a5c3e50 --- /dev/null +++ b/internal/plugins/telephony/telephony_test.go @@ -0,0 +1,80 @@ +package telephony + +import ( + "context" + "crypto/x509" + "encoding/json" + "net" + "testing" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap/zaptest" +) + +type stubSender struct{} + +func (stubSender) ID() string { return "dev1" } +func (stubSender) Name() string { return "test" } +func (stubSender) SetName(string) {} +func (stubSender) State() device.PairingState { return device.StatePaired } +func (stubSender) SetState(device.PairingState) {} +func (stubSender) Send(*protocol.Packet) error { return nil } +func (stubSender) IsConnected() bool { return true } +func (stubSender) RemoteIP() net.IP { return nil } +func (stubSender) PeerCert() *x509.Certificate { return nil } +func (stubSender) HasCapability(string) bool { return true } +func (stubSender) UpdateBattery(int, bool) {} +func (stubSender) GetBattery() (int, bool) { return 0, false } + +func nextEvent(t *testing.T, ch <-chan events.Event) events.Event { + t.Helper() + select { + case ev := <-ch: + return ev + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for event") + return events.Event{} + } +} + +func TestHandleStringIsCancel(t *testing.T) { + // Stock phones end calls with the string "true", not a boolean. + logger := zaptest.NewLogger(t) + bus := events.NewBus(logger) + p := NewTelephonyPlugin(bus, logger) + sub := bus.Subscribe(0, events.TypeTelephonyCanceled) + defer sub.Close() + + pkt := &protocol.Packet{ + Type: "kdeconnect.telephony", + Body: json.RawMessage(`{"event":"ringing","phoneNumber":"+1","isCancel":"true"}`), + } + if err := p.Handle(context.Background(), stubSender{}, pkt); err != nil { + t.Fatalf("Handle returned error: %v", err) + } + if ev := nextEvent(t, sub.C); ev.Type != events.TypeTelephonyCanceled { + t.Errorf("event type = %q, want %q", ev.Type, events.TypeTelephonyCanceled) + } +} + +func TestHandleMissedCallEvent(t *testing.T) { + logger := zaptest.NewLogger(t) + bus := events.NewBus(logger) + p := NewTelephonyPlugin(bus, logger) + sub := bus.Subscribe(0, events.TypeTelephonyMissed) + defer sub.Close() + + pkt := &protocol.Packet{ + Type: "kdeconnect.telephony", + Body: json.RawMessage(`{"event":"missedCall","contactName":"Alice"}`), + } + if err := p.Handle(context.Background(), stubSender{}, pkt); err != nil { + t.Fatalf("Handle returned error: %v", err) + } + if ev := nextEvent(t, sub.C); ev.Type != events.TypeTelephonyMissed { + t.Errorf("event type = %q, want %q", ev.Type, events.TypeTelephonyMissed) + } +} From 8a72611d41299bd60565abae9bce35d349ad6c1d Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:58:36 +0300 Subject: [PATCH 07/12] feat(notification): clear phone notifications from desktop Dismissals previously had no send path, so clearing a notification here left it standing on the phone. A new notify_dismiss command flows from the dismiss CLI through the daemon to a cancel request packet, and the matching desktop popup closes with it. --- cmd/kcd/cli_dismiss.go | 27 +++++++++++ cmd/kcd/main.go | 1 + docs/CLI.md | 16 +++++++ docs/CLIENT_GUIDE.md | 8 ++++ docs/IPC_PROTOCOL.md | 18 ++++++- internal/daemon/ipc_routes_comm.go | 18 +++++++ internal/ipc/proto.go | 7 +++ internal/plugins/notification/notification.go | 24 +++++++++- .../plugins/notification/notification_test.go | 47 +++++++++++++++++++ pkg/client/client.go | 9 ++++ 10 files changed, 173 insertions(+), 2 deletions(-) create mode 100644 cmd/kcd/cli_dismiss.go diff --git a/cmd/kcd/cli_dismiss.go b/cmd/kcd/cli_dismiss.go new file mode 100644 index 0000000..edc1933 --- /dev/null +++ b/cmd/kcd/cli_dismiss.go @@ -0,0 +1,27 @@ +package main + +import ( + "fmt" + + "github.com/urfave/cli/v2" +) + +var dismissCmd = &cli.Command{ + Name: "dismiss", + Usage: "Clear a notification on the smartphone", + ArgsUsage: " ", + Action: func(c *cli.Context) error { + if c.NArg() < 2 { + return fmt.Errorf("missing device ID or notification ID") + } + cl, err := getClient(c) + if err != nil { + return err + } + if err := cl.NotifyDismiss(c.Args().Get(0), c.Args().Get(1)); err != nil { + return err + } + fmt.Println("Notification dismissed") + return nil + }, +} diff --git a/cmd/kcd/main.go b/cmd/kcd/main.go index 54a7219..5443d38 100644 --- a/cmd/kcd/main.go +++ b/cmd/kcd/main.go @@ -130,6 +130,7 @@ func main() { watchCmd, sftpCmd, replyCmd, + dismissCmd, callCmd, findmyphoneCmd, lockCmd, diff --git a/docs/CLI.md b/docs/CLI.md index 151a295..00e494a 100644 --- a/docs/CLI.md +++ b/docs/CLI.md @@ -582,6 +582,22 @@ kcd reply a1b2... abc-123 "On my way!" --- +## dismiss + +Clear a notification on the phone and close its desktop popup. + +``` +kcd dismiss +``` + +The `notification-id` is the `id` field of a `notification` event: + +```bash +kcd dismiss a1b2... notif-456 +``` + +--- + ## call Manage phone calls. diff --git a/docs/CLIENT_GUIDE.md b/docs/CLIENT_GUIDE.md index efee2ec..c2249cf 100644 --- a/docs/CLIENT_GUIDE.md +++ b/docs/CLIENT_GUIDE.md @@ -471,6 +471,8 @@ Each event type carries a different payload shape. Here are the common ones: - `requestReplyId` is present only for notifications that support inline replies. Use it with the `notify_reply` command. +- `id` identifies the notification for `notify_dismiss`, which clears it + on the phone and closes the desktop popup. ### Share Progress @@ -574,6 +576,12 @@ def reply_to_notification(dev_id, notif_id, message): "replyId": notif_id, "message": message }) + +def dismiss_notification(dev_id, notif_id): + ipc_request(command_sock, "notify_dismiss", { + "deviceId": dev_id, + "notificationId": notif_id + }) ``` ### 7.3 SMS Viewer diff --git a/docs/IPC_PROTOCOL.md b/docs/IPC_PROTOCOL.md index aa52927..d92191d 100644 --- a/docs/IPC_PROTOCOL.md +++ b/docs/IPC_PROTOCOL.md @@ -416,6 +416,21 @@ The `replyId` comes from the `requestReplyId` field of a `notification` event. **Response data:** none +#### `notify_dismiss` + +Clear a notification on the phone (sends `kdeconnect.notification.request` +with `{"cancel": ""}`) and close the matching desktop popup. + +**Request payload:** + +```json +{"deviceId": "a1b2c3d4e5f6_...", "notificationId": "notif-456"} +``` + +The `notificationId` is the `id` field of a `notification` event. + +**Response data:** none + #### `findmyphone` (also aliased as `ring`) Make a paired phone ring loudly. @@ -1287,6 +1302,7 @@ who may want to implement a full network-level implementation. | `kdeconnect.mpris` | MPRIS | Player list, NowPlaying state, seek positions, album art (broadcast + request-reply) | | `kdeconnect.mpris.request` | MPRIS | Request player list, now-playing, volume, album art; send control actions | | `kdeconnect.notification.reply` | Notification | Reply to a notification with inline reply support | +| `kdeconnect.notification.request` | Notification | Clear a notification on the phone (`{"cancel": ""}`) | | `kdeconnect.notification` | RunCommand | Command output notification pushed to phone | | `kdeconnect.runcommand` | RunCommand | Send command list to phone | | `kdeconnect.runcommand.request` | RunCommand | Request phone's command list / execute command | @@ -1315,7 +1331,7 @@ plugin processes it and a link to the body struct definition. | Packet Type | Plugin | Body Struct | |---|---|---| | `kdeconnect.pair` | Pair | `PairBody{Pair bool, Timestamp int64}` | -| `kdeconnect.battery` | Battery | `BatteryBody{CurrentCharge int, IsCharging bool, ThresholdEvent int}` | +| `kdeconnect.battery` | Battery | `BatteryBody{CurrentCharge int, IsCharging bool, ThresholdEvent int, Request bool}` — `request:true` asks for state, never stored | | `kdeconnect.battery.request` | Battery | (empty, triggers a battery reply) | | `kdeconnect.notification` | Notification | `NotificationBody{ID, AppName, Title, Text, IsCancel, IsClearable, Silent, RequestReplyId string}` | | `kdeconnect.share.request` | Share | `ShareBody{Filename, NumberOfFiles, TotalPayloadSize, LastModified, CreationTime, Text, Url}` | diff --git a/internal/daemon/ipc_routes_comm.go b/internal/daemon/ipc_routes_comm.go index 364e8ee..83f4a79 100644 --- a/internal/daemon/ipc_routes_comm.go +++ b/internal/daemon/ipc_routes_comm.go @@ -129,5 +129,23 @@ func registerCommRoutes(handler *ipc.Handler, cfg *config.Config, devices *devic } return ipc.Response{OK: true} }) + handler.Register(ipc.CmdNotifyDismiss, func(req ipc.Request) ipc.Response { + var p ipc.NotifyDismissPayload + if err := json.Unmarshal(req.Payload, &p); err != nil { + return ipc.Response{OK: false, Error: "invalid payload"} + } + pl, ok := plugins.GetByName("Notification") + if !ok { + return ipc.Response{OK: false, Error: "notification plugin not enabled"} + } + dev, ok := devices.Get(p.DeviceID) + if !ok { + return ipc.Response{OK: false, Error: "device not found"} + } + if err := pl.(*notification.NotificationPlugin).Dismiss(dev, p.NotificationID); err != nil { + return ipc.Response{OK: false, Error: err.Error()} + } + return ipc.Response{OK: true} + }) } } diff --git a/internal/ipc/proto.go b/internal/ipc/proto.go index e0adfb2..4dbcd12 100644 --- a/internal/ipc/proto.go +++ b/internal/ipc/proto.go @@ -25,6 +25,7 @@ const ( CmdSftpInfo = "sftp_info" CmdSftpVolumes = "sftp_volumes" CmdNotifyReply = "notify_reply" + CmdNotifyDismiss = "notify_dismiss" CmdCallMute = "call_mute" CmdFindMyPhone = "findmyphone" CmdLock = "lock" @@ -99,6 +100,12 @@ type NotifyReplyPayload struct { Message string `json:"message"` } +// NotifyDismissPayload is used for CmdNotifyDismiss. +type NotifyDismissPayload struct { + DeviceID string `json:"deviceId"` + NotificationID string `json:"notificationId"` +} + // SMSPayload is used for CmdSendSMS. type SMSPayload struct { DeviceID string `json:"deviceId"` diff --git a/internal/plugins/notification/notification.go b/internal/plugins/notification/notification.go index b768f1a..f9d416b 100644 --- a/internal/plugins/notification/notification.go +++ b/internal/plugins/notification/notification.go @@ -129,7 +129,7 @@ func (p *NotificationPlugin) IncomingTypes() []string { return []string{"kdeconnect.notification"} } func (p *NotificationPlugin) OutgoingTypes() []string { - return []string{"kdeconnect.notification.reply"} + return []string{"kdeconnect.notification.reply", "kdeconnect.notification.request"} } // nonAlphaNumeric sanitises app names to be safe for exec / notify-send args. @@ -437,5 +437,27 @@ func (p *NotificationPlugin) RequestReply(dev device.Sender, replyID, message st return dev.Send(pkt) } +// Dismiss asks the phone to clear the notification with the given ID and +// closes the matching desktop popup, if one is tracked. +func (p *NotificationPlugin) Dismiss(dev device.Sender, id string) error { + pkt, err := protocol.NewPacket("kdeconnect.notification.request", map[string]string{ + "cancel": id, + }) + if err != nil { + return err + } + if err := dev.Send(pkt); err != nil { + return err + } + if id != "" { + if desktopID, ok := p.notifIDs.LoadAndDelete(p.notifKey(dev.ID(), id)); ok { + if s, ok := desktopID.(string); ok { + p.closeNotification(s) + } + } + } + return nil +} + func (p *NotificationPlugin) OnConnect(_ device.Sender) {} func (p *NotificationPlugin) OnDisconnect(_ device.Sender) {} diff --git a/internal/plugins/notification/notification_test.go b/internal/plugins/notification/notification_test.go index 8f66d11..4813433 100644 --- a/internal/plugins/notification/notification_test.go +++ b/internal/plugins/notification/notification_test.go @@ -3,7 +3,10 @@ package notification import ( "context" "crypto/tls" + "crypto/x509" + "encoding/json" "fmt" + "net" "os" "os/exec" "path/filepath" @@ -449,3 +452,47 @@ func TestNotificationPlugin_CancelGraceDisabled(t *testing.T) { t.Fatalf("expected gdbus close of desktop id 1, got %q", got) } } + +// captureSender implements device.Sender and records outbound packets. +type captureSender struct { + id string + sent []*protocol.Packet +} + +func (s *captureSender) ID() string { return s.id } +func (s *captureSender) Name() string { return "test" } +func (s *captureSender) SetName(string) {} +func (s *captureSender) State() device.PairingState { return device.StatePaired } +func (s *captureSender) SetState(device.PairingState) {} +func (s *captureSender) Send(p *protocol.Packet) error { s.sent = append(s.sent, p); return nil } +func (s *captureSender) IsConnected() bool { return true } +func (s *captureSender) RemoteIP() net.IP { return nil } +func (s *captureSender) PeerCert() *x509.Certificate { return nil } +func (s *captureSender) HasCapability(string) bool { return true } +func (s *captureSender) UpdateBattery(int, bool) {} +func (s *captureSender) GetBattery() (int, bool) { return 0, false } + +func TestNotificationPlugin_Dismiss(t *testing.T) { + p := newPlugin(t) + dev := &captureSender{id: "dev1"} + + if err := p.Dismiss(dev, "notif-123"); err != nil { + t.Fatalf("Dismiss returned error: %v", err) + } + if len(dev.sent) != 1 { + t.Fatalf("sent %d packets, want 1", len(dev.sent)) + } + pkt := dev.sent[0] + if pkt.Type != "kdeconnect.notification.request" { + t.Errorf("packet type = %q, want kdeconnect.notification.request", pkt.Type) + } + var body struct { + Cancel string `json:"cancel"` + } + if err := json.Unmarshal(pkt.Body, &body); err != nil { + t.Fatalf("unmarshal dismiss body: %v", err) + } + if body.Cancel != "notif-123" { + t.Errorf("cancel = %q, want %q", body.Cancel, "notif-123") + } +} diff --git a/pkg/client/client.go b/pkg/client/client.go index 733eebb..2d8a5ea 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -298,6 +298,15 @@ func (c *Client) NotifyReply(deviceID, replyID, message string) error { return err } +// NotifyDismiss requests the daemon to clear a notification on the remote device. +func (c *Client) NotifyDismiss(deviceID, notificationID string) error { + _, err := c.Call(ipc.CmdNotifyDismiss, ipc.NotifyDismissPayload{ + DeviceID: deviceID, + NotificationID: notificationID, + }) + return err +} + // CallMute requests the daemon to mute an incoming call on the remote device. func (c *Client) CallMute(deviceID string) error { _, err := c.Call(ipc.CmdCallMute, ipc.DevicePayload{DeviceID: deviceID}) From faa19b4333ebfc8d546d253c502eab5ee1b127f5 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 04:02:16 +0300 Subject: [PATCH 08/12] fix(transfer): abort side-channel transfers on sustained silence MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Transfers streamed with no deadline at all, so a phone leaving Wi-Fi mid-file left io.Copy blocked until TCP keepalive gave up minutes later. Streams now carry an activity deadline that renews on every read and write — slow links run as long as they make progress, dead ones fail fast — governed by network.transfer_idle_timeout, with corrupt partials removed on failure. --- docs/CLI.md | 2 +- internal/config/config.go | 2 +- internal/config/config_test.go | 5 +- internal/config/runtime.go | 10 ++-- internal/daemon/plugins.go | 3 +- internal/plugins/clipboard/clipboard.go | 1 + internal/plugins/share/share.go | 2 +- internal/plugins/share/transfer.go | 12 ++++- internal/plugins/sms/sms.go | 1 + internal/transport/idleconn.go | 41 +++++++++++++++ internal/transport/idleconn_test.go | 70 +++++++++++++++++++++++++ internal/transport/sidechannel.go | 17 +++--- packaging/kcd.example.toml | 1 + 13 files changed, 149 insertions(+), 18 deletions(-) create mode 100644 internal/transport/idleconn.go create mode 100644 internal/transport/idleconn_test.go diff --git a/docs/CLI.md b/docs/CLI.md index 00e494a..aa95d63 100644 --- a/docs/CLI.md +++ b/docs/CLI.md @@ -29,7 +29,7 @@ require a restart; reloading notification filters alone does not apply them. | Section | Settings and defaults | |---|---| -| `[network]` | `dial_timeout = "5s"`, `handshake_timeout = "10s"`, `sidechannel_timeout = "15s"` | +| `[network]` | `dial_timeout = "5s"`, `handshake_timeout = "10s"`, `sidechannel_timeout = "15s"`, `transfer_idle_timeout = "60s"` | | `[reconnect]` | `initial_backoff = "2s"`, `max_backoff = "5m"`, `flap_threshold = "15s"` | | `[discovery]` | `broadcast_interval = "30s"`, `broadcast_idle_interval = "60s"` | | `[pairing]` | `intent_ttl = "5m"`, `listen_timeout = "60s"`; existing `timeout_secs = 30` still controls the pairing response wait | diff --git a/internal/config/config.go b/internal/config/config.go index 4de20cc..be0f80a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -69,7 +69,7 @@ func Defaults() *Config { c.TCPPort = 1716 c.LogLevel = "info" - c.Network = NetworkConfig{DialTimeout: "5s", HandshakeTimeout: "10s", SidechannelTimeout: "15s"} + c.Network = NetworkConfig{DialTimeout: "5s", HandshakeTimeout: "10s", SidechannelTimeout: "15s", TransferIdleTimeout: "60s"} c.Reconnect = ReconnectConfig{InitialBackoff: "2s", MaxBackoff: "5m", FlapThreshold: "15s"} c.Discovery = DiscoveryConfig{BroadcastInterval: "30s", BroadcastIdleInterval: "60s"} c.Plugins.Defaults() diff --git a/internal/config/config_test.go b/internal/config/config_test.go index eb024da..3139633 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -16,7 +16,7 @@ func TestDefaults(t *testing.T) { if err := cfg.Validate(); err != nil { t.Fatal(err) } - if cfg.Network != (NetworkConfig{"5s", "10s", "15s"}) { + if cfg.Network != (NetworkConfig{"5s", "10s", "15s", "60s"}) { t.Errorf("network defaults: %+v", cfg.Network) } if cfg.Reconnect != (ReconnectConfig{"2s", "5m", "15s"}) { @@ -68,6 +68,7 @@ func TestLoadOverridesAndRoundTrip(t *testing.T) { dial_timeout = "750ms" handshake_timeout = "12s" sidechannel_timeout = "25s" +transfer_idle_timeout = "90s" [reconnect] initial_backoff = "3s" max_backoff = "6m" @@ -92,7 +93,7 @@ app_name = "KDE Connect" if err != nil { t.Fatal(err) } - if cfg.Network != (NetworkConfig{"750ms", "12s", "25s"}) || cfg.Reconnect != (ReconnectConfig{"3s", "6m", "20s"}) || cfg.Discovery != (DiscoveryConfig{"45s", "90s"}) { + if cfg.Network != (NetworkConfig{"750ms", "12s", "25s", "90s"}) || cfg.Reconnect != (ReconnectConfig{"3s", "6m", "20s"}) || cfg.Discovery != (DiscoveryConfig{"45s", "90s"}) { t.Fatal("duration overrides not decoded") } if cfg.Cache != (CacheConfig{"/tmp/sms", "/tmp/art", "/tmp/contacts"}) || cfg.Pairing.IntentTTL != "7m" || cfg.Pairing.ListenTimeout != "2m" { diff --git a/internal/config/runtime.go b/internal/config/runtime.go index acff8d2..d3b0495 100644 --- a/internal/config/runtime.go +++ b/internal/config/runtime.go @@ -6,10 +6,13 @@ import ( ) // NetworkConfig controls connection setup deadlines, not transfer size limits. +// TransferIdleTimeout bounds streaming silence on side-channel transfers: +// any read/write gap longer than it aborts the transfer. type NetworkConfig struct { - DialTimeout string `toml:"dial_timeout"` - HandshakeTimeout string `toml:"handshake_timeout"` - SidechannelTimeout string `toml:"sidechannel_timeout"` + DialTimeout string `toml:"dial_timeout"` + HandshakeTimeout string `toml:"handshake_timeout"` + SidechannelTimeout string `toml:"sidechannel_timeout"` + TransferIdleTimeout string `toml:"transfer_idle_timeout"` } // ReconnectConfig controls retry delays and the minimum stable connection age. @@ -47,6 +50,7 @@ func (c *Config) validateDurations() error { {"network.dial_timeout", c.Network.DialTimeout}, {"network.handshake_timeout", c.Network.HandshakeTimeout}, {"network.sidechannel_timeout", c.Network.SidechannelTimeout}, + {"network.transfer_idle_timeout", c.Network.TransferIdleTimeout}, {"reconnect.initial_backoff", c.Reconnect.InitialBackoff}, {"reconnect.max_backoff", c.Reconnect.MaxBackoff}, {"reconnect.flap_threshold", c.Reconnect.FlapThreshold}, diff --git a/internal/daemon/plugins.go b/internal/daemon/plugins.go index 6c46ba8..e375ebc 100644 --- a/internal/daemon/plugins.go +++ b/internal/daemon/plugins.go @@ -34,7 +34,8 @@ import ( func setupPlugins(cfg *config.Config, bus *events.Bus, tlsCfg *tls.Config, logger *zap.Logger, devices *device.Registry, localCert *x509.Certificate, saveDevices func(), plugins *plugin.Registry) *pair.PairPlugin { sidechannel := transport.SidechannelOptions{ - Timeout: config.Duration(cfg.Network.SidechannelTimeout), + Timeout: config.Duration(cfg.Network.SidechannelTimeout), + IdleTimeout: config.Duration(cfg.Network.TransferIdleTimeout), } pairPlugin := pair.NewPairPlugin(devices, localCert, cfg.Pairing, saveDevices, bus, logger) plugins.Register(pairPlugin) diff --git a/internal/plugins/clipboard/clipboard.go b/internal/plugins/clipboard/clipboard.go index d8a908a..b894f57 100644 --- a/internal/plugins/clipboard/clipboard.go +++ b/internal/plugins/clipboard/clipboard.go @@ -393,6 +393,7 @@ func downloadToFile(ctx context.Context, ip net.IP, port int, size int64, dest s _, err = io.Copy(f, io.LimitReader(conn, size)) if err != nil { + os.Remove(dest) // don't leave a corrupt partial behind return fmt.Errorf("clipboard: stream to %s: %w", dest, err) } return nil diff --git a/internal/plugins/share/share.go b/internal/plugins/share/share.go index 4c4da88..3c0da48 100644 --- a/internal/plugins/share/share.go +++ b/internal/plugins/share/share.go @@ -276,7 +276,7 @@ func (p *SharePlugin) SendFile(ctx context.Context, dev device.Sender, filePath if timeout == 0 { timeout = 2 * time.Minute } - err := AcceptAndSend(ln, filePath, p.TLSConfig, dev.ID(), expectedFP, timeout, onProgress, p.Logger) + err := AcceptAndSend(ln, filePath, p.TLSConfig, dev.ID(), expectedFP, timeout, onProgress, p.Logger, p.sidechannel) if err != nil { p.Logger.Error("share: send failed", diff --git a/internal/plugins/share/transfer.go b/internal/plugins/share/transfer.go index e4bdb1a..b028311 100644 --- a/internal/plugins/share/transfer.go +++ b/internal/plugins/share/transfer.go @@ -57,10 +57,12 @@ func ReceiveSideChannel(ctx context.Context, ip net.IP, port int, size int64, de n, err := io.Copy(f, r) if err != nil { + os.Remove(dest) // don't leave a corrupt partial behind return fmt.Errorf("share: stream transfer to %s: %w", dest, err) } if n < size { + os.Remove(dest) // don't leave a corrupt partial behind return fmt.Errorf("share: transfer truncated (%d/%d bytes)", n, size) } @@ -84,7 +86,8 @@ func ListenSideChannel(ctx context.Context, cfg config.ShareConfig, tlsConfig *t } // AcceptAndSend waits for the phone to connect, performs the TLS handshake, and streams the file. -func AcceptAndSend(ln net.Listener, filePath string, tlsConfig *tls.Config, expectedDeviceID, expectedFP string, timeout time.Duration, onProgress func(int64, int64), logger *zap.Logger) error { +// Optional SidechannelOptions bound streaming silence the same way as the dial path. +func AcceptAndSend(ln net.Listener, filePath string, tlsConfig *tls.Config, expectedDeviceID, expectedFP string, timeout time.Duration, onProgress func(int64, int64), logger *zap.Logger, options ...transport.SidechannelOptions) error { defer ln.Close() addr := ln.Addr().String() @@ -190,7 +193,12 @@ func AcceptAndSend(ln net.Listener, filePath string, tlsConfig *tls.Config, expe r = io.TeeReader(f, &progressWriter{total: size, callback: onProgress}) } - n, err := io.Copy(tlsConn, r) + var streamConn net.Conn = tlsConn + if len(options) > 0 { + streamConn = transport.WithIdleTimeout(tlsConn, options[0].IdleTimeout) + } + + n, err := io.Copy(streamConn, r) if err != nil { return fmt.Errorf("stream error: %w", err) } diff --git a/internal/plugins/sms/sms.go b/internal/plugins/sms/sms.go index 8b8b86b..a6956a8 100644 --- a/internal/plugins/sms/sms.go +++ b/internal/plugins/sms/sms.go @@ -289,6 +289,7 @@ func (p *SMSPlugin) receiveAttachment(ctx context.Context, ip net.IP, port int, _, err = io.Copy(f, io.LimitReader(conn, size)) if err != nil { + os.Remove(destPath) // don't leave a corrupt partial behind return fmt.Errorf("sms: receive attachment data: %w", err) } diff --git a/internal/transport/idleconn.go b/internal/transport/idleconn.go new file mode 100644 index 0000000..1693e5e --- /dev/null +++ b/internal/transport/idleconn.go @@ -0,0 +1,41 @@ +package transport + +import ( + "net" + "time" +) + +// IdleTimeoutConn bounds silence, not size: every Read or Write pushes its +// deadline forward by the idle interval, so slow-but-alive transfers run as +// long as they need while a stalled socket (walked-out-of-range phone, +// half-open zombie) fails promptly instead of hanging on TCP keepalive. +// Wrap with WithIdleTimeout; a non-positive interval returns conn unchanged. +type IdleTimeoutConn struct { + net.Conn + idle time.Duration +} + +// WithIdleTimeout wraps conn so any read or write gap longer than idle +// fails with a timeout error. +func WithIdleTimeout(conn net.Conn, idle time.Duration) net.Conn { + if conn == nil || idle <= 0 { + return conn + } + return &IdleTimeoutConn{Conn: conn, idle: idle} +} + +// Read sets a fresh read deadline before every read. +func (c *IdleTimeoutConn) Read(b []byte) (int, error) { + if err := c.SetReadDeadline(time.Now().Add(c.idle)); err != nil { + return 0, err + } + return c.Conn.Read(b) +} + +// Write sets a fresh write deadline before every write. +func (c *IdleTimeoutConn) Write(b []byte) (int, error) { + if err := c.SetWriteDeadline(time.Now().Add(c.idle)); err != nil { + return 0, err + } + return c.Conn.Write(b) +} diff --git a/internal/transport/idleconn_test.go b/internal/transport/idleconn_test.go new file mode 100644 index 0000000..1c5c5c7 --- /dev/null +++ b/internal/transport/idleconn_test.go @@ -0,0 +1,70 @@ +package transport + +import ( + "errors" + "net" + "testing" + "time" +) + +func TestWithIdleTimeoutDisabled(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + if got := WithIdleTimeout(a, 0); got != a { + t.Error("zero idle must return the conn unwrapped") + } + if got := WithIdleTimeout(nil, time.Second); got != nil { + t.Error("nil conn must stay nil") + } + _ = b +} + +func TestIdleTimeoutFiresOnSilence(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + r := WithIdleTimeout(a, 50*time.Millisecond) + buf := make([]byte, 8) + start := time.Now() + _, err := r.Read(buf) + if err == nil { + t.Fatal("silent read succeeded, want timeout") + } + var netErr net.Error + if !errors.As(err, &netErr) || !netErr.Timeout() { + t.Fatalf("error = %v, want timeout", err) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("read blocked %v, want prompt timeout", elapsed) + } +} + +func TestIdleTimeoutExtendedByActivity(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + r := WithIdleTimeout(a, 100*time.Millisecond) + done := make(chan error, 1) + go func() { + // Drip-feed bytes slower in total than the idle window but faster + // per-gap: total ~150ms of streaming must survive a 100ms idle. + for i := 0; i < 5; i++ { + time.Sleep(30 * time.Millisecond) + if _, err := b.Write([]byte("x")); err != nil { + done <- err + return + } + } + done <- nil + }() + buf := make([]byte, 5) + for i := 0; i < 5; i++ { + if _, err := r.Read(buf[i : i+1]); err != nil { + t.Fatalf("active read %d failed: %v", i, err) + } + } + if err := <-done; err != nil { + t.Fatalf("writer failed: %v", err) + } +} diff --git a/internal/transport/sidechannel.go b/internal/transport/sidechannel.go index f8dd766..52f330f 100644 --- a/internal/transport/sidechannel.go +++ b/internal/transport/sidechannel.go @@ -12,17 +12,20 @@ import ( "go.uber.org/zap" ) -// SidechannelOptions bounds the entire connection establishment phase — TCP -// dial plus TLS handshake plus pin verification — with a single overall -// timeout (default 15s). Payload streaming is never time-bounded. +// SidechannelOptions bounds the connection establishment phase — TCP dial +// plus TLS handshake plus pin verification — with a single overall timeout +// (default 15s). IdleTimeout optionally bounds streaming silence: every +// read or write pushes the deadline forward, so stalled transfers fail +// while slow-but-alive ones run unbounded in total. Zero disables it. type SidechannelOptions struct { - Timeout time.Duration + Timeout time.Duration + IdleTimeout time.Duration } // DialSidechannel connects as a TLS client and verifies the paired certificate // before returning any payload bytes. The caller owns and must close the result. -// Once established, payload streaming carries no deadline, so large transfers -// are never truncated mid-stream by the helper. +// Once established, payload streaming carries no absolute deadline — only the +// optional idle bound from SidechannelOptions applies. func DialSidechannel(ctx context.Context, ip net.IP, port int, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger, options ...SidechannelOptions) (net.Conn, error) { if ip == nil || port < 1 || port > 65535 { return nil, fmt.Errorf("side-channel: invalid peer address") @@ -66,5 +69,5 @@ func DialSidechannel(ctx context.Context, ip net.IP, port int, tlsConfig *tls.Co // unbounded, and download goroutines may legitimately outlive the // packet-handler context. Callers own the returned conn and close it // when streaming finishes. - return conn, nil + return WithIdleTimeout(conn, opts.IdleTimeout), nil } diff --git a/packaging/kcd.example.toml b/packaging/kcd.example.toml index 9b22289..1a22b5d 100644 --- a/packaging/kcd.example.toml +++ b/packaging/kcd.example.toml @@ -171,6 +171,7 @@ remotesystemvolume = true # dial_timeout = "5s" # control-channel TCP dial # handshake_timeout = "10s" # identity/TLS handshake # sidechannel_timeout = "15s" # side-channel connection setup +# transfer_idle_timeout = "60s" # abort transfers silent longer than this [reconnect] # initial_backoff = "2s" From 400a8b9b7e524d83b7a89be4c64107a8e6fa7761 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 04:15:43 +0300 Subject: [PATCH 09/12] fix(connect): leave manual dials unaddressed so phones accept them Connect-by-IP stamped the placeholder "manual" into the pre-TLS identity's targetDeviceId, and stock phones drop any pre-TLS identity addressed to another device. Unknown targets now omit the field entirely, which peers treat as broadcast, and the dial uses the configured TCP port instead of a hardcoded one. --- internal/daemon/ipc_routes.go | 5 ++- internal/daemon/transport.go | 3 ++ internal/daemon/transport_test.go | 70 +++++++++++++++++++++++++++++++ 3 files changed, 77 insertions(+), 1 deletion(-) diff --git a/internal/daemon/ipc_routes.go b/internal/daemon/ipc_routes.go index 581348b..dbb54ae 100644 --- a/internal/daemon/ipc_routes.go +++ b/internal/daemon/ipc_routes.go @@ -72,7 +72,10 @@ func registerIPCRoutes(handler *ipc.Handler, cfg *config.Config, devices *device if err != nil { return } - DialDevice(ctx, addr, 1716, "manual", protocol.ProtocolVersion, identityPkt, tlsCfg, devices, plugins, cfg.DeviceID, logger, true, cfg) + // No target ID: the peer is whoever answers at this address. + // An empty target omits targetDeviceId from the pre-TLS + // identity; stock peers drop dials addressed to anyone else. + DialDevice(ctx, addr, cfg.TCPPort, "", protocol.ProtocolVersion, identityPkt, tlsCfg, devices, plugins, cfg.DeviceID, logger, true, cfg) }() return ipc.Response{OK: true} diff --git a/internal/daemon/transport.go b/internal/daemon/transport.go index d525e30..16fbcd9 100644 --- a/internal/daemon/transport.go +++ b/internal/daemon/transport.go @@ -76,6 +76,9 @@ func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID s if targetProto <= 0 || targetProto > protocol.ProtocolVersion { targetProto = protocol.ProtocolVersion } + // TargetDeviceID stays empty (and thus absent on the wire) when the + // target is unknown, e.g. an explicit connect-by-IP: stock peers + // close pre-TLS identities addressed to any other device ID. preTlsId := protocol.IdentityBody{ DeviceID: myID.DeviceID, DeviceName: myID.DeviceName, diff --git a/internal/daemon/transport_test.go b/internal/daemon/transport_test.go index f03487d..b7aa924 100644 --- a/internal/daemon/transport_test.go +++ b/internal/daemon/transport_test.go @@ -2,12 +2,18 @@ package daemon import ( "context" + "crypto/tls" + "encoding/json" + "net" "testing" "time" + "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/discovery" + "github.com/bethropolis/kcd/internal/plugin" "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" "go.uber.org/zap" ) @@ -191,3 +197,67 @@ func TestSyncReconnectBroadcast(t *testing.T) { t.Error("pairing stop must end an unowned loop") } } + +// A manual connect-by-IP dial must not address the pre-TLS identity to any +// device: stock peers close identities carrying a foreign targetDeviceId, +// so the field has to stay absent. Dials with a known target keep it. +func TestDialPreTLSIdentityTarget(t *testing.T) { + cases := []struct { + name string + targetID string + wantTarget string // empty means the key must be absent + }{ + {"manual dial omits target", "", ""}, + {"known target kept", "peer-device-id", "peer-device-id"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + ln, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + got := make(chan []byte, 1) + go func() { + c, err := ln.Accept() + if err != nil { + return + } + defer c.Close() + pkt, _, err := transport.ReadPlaintextPacket(c) + if err != nil { + return + } + defer protocol.ReleasePacket(pkt) + got <- append([]byte(nil), pkt.Body...) + }() + cfg := config.Defaults() + pkt, err := protocol.NewIdentityPacket("local", "Local", "desktop", cfg.TCPPort, nil, nil) + if err != nil { + t.Fatal(err) + } + logger := zap.NewNop() + addr := ln.Addr().(*net.TCPAddr) + DialDevice(context.Background(), addr.IP, addr.Port, tc.targetID, protocol.ProtocolVersion, pkt, &tls.Config{}, device.NewRegistry(nil), plugin.NewRegistry(logger), "local", logger, true, cfg) + select { + case body := <-got: + var fields map[string]any + if err := json.Unmarshal(body, &fields); err != nil { + t.Fatalf("unmarshal pre-TLS body: %v", err) + } + v, present := fields["targetDeviceId"] + if tc.wantTarget == "" { + if present { + t.Errorf("targetDeviceId present with %v, want absent", v) + } + return + } + if !present || v != tc.wantTarget { + t.Errorf("targetDeviceId = %v, want %q", v, tc.wantTarget) + } + case <-time.After(10 * time.Second): + t.Fatal("timed out waiting for pre-TLS identity") + } + }) + } +} From 67e3d5574c35cbc4e7873ddcc2e17e4d27b404e9 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Thu, 17 Sep 2026 04:34:04 +0300 Subject: [PATCH 10/12] feat(status): per-device table with link, battery and last-seen The status summary knew only counts, forcing a second command for anything about a device. The daemon now reports each known device with its state, address, cached battery and last-seen time, and the CLI renders it as an aligned, sectioned table with the plugin inventory counted. --- cmd/kcd/cli_status.go | 87 +++++++++++++++++++++++++++++++++++ cmd/kcd/cli_status_test.go | 61 ++++++++++++++++++++++++ cmd/kcd/main.go | 7 +-- docs/CLI.md | 17 +++++-- internal/daemon/ipc_routes.go | 32 ++++++++++++- internal/ipc/proto.go | 37 +++++++++++---- 6 files changed, 221 insertions(+), 20 deletions(-) create mode 100644 cmd/kcd/cli_status.go create mode 100644 cmd/kcd/cli_status_test.go diff --git a/cmd/kcd/cli_status.go b/cmd/kcd/cli_status.go new file mode 100644 index 0000000..95725b4 --- /dev/null +++ b/cmd/kcd/cli_status.go @@ -0,0 +1,87 @@ +package main + +import ( + "fmt" + "strings" + "text/tabwriter" + "time" + + "github.com/bethropolis/kcd/internal/ipc" +) + +// formatStatus renders the human-readable `kcd status` output. +func formatStatus(st *ipc.StatusResponse) string { + var b strings.Builder + fmt.Fprintf(&b, "kcd %s (up %s)\n", st.Version, st.UptimeHuman) + fmt.Fprintf(&b, "\nSocket: %s\n", st.SocketPath) + fmt.Fprintf(&b, "Config: %s\n", st.ConfigPath) + if st.TCPPort > 0 { + fmt.Fprintf(&b, "Listen: tcp :%d\n", st.TCPPort) + } + fmt.Fprintf(&b, "\nDevices: %d known, %d connected\n", st.DeviceCount, st.ConnectedCount) + if len(st.Devices) > 0 { + w := tabwriter.NewWriter(&b, 0, 4, 2, ' ', 0) + fmt.Fprintln(w, "NAME\tID\tTYPE\tSTATE\tADDR\tBATTERY\tLAST SEEN") + for _, d := range st.Devices { + fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\t%s\n", + d.Name, shortDeviceID(d.ID), d.Type, d.State, + orDash(d.Addr), formatStatusBattery(d.Battery), formatSeenAge(d.LastSeen)) + } + w.Flush() + } + fmt.Fprintf(&b, "\nPlugins (%d): %s\n", len(st.Plugins), strings.Join(st.Plugins, ", ")) + return b.String() +} + +// shortDeviceID shows the leading segment of a device ID. +func shortDeviceID(id string) string { + if i := strings.Index(id, "_"); i > 0 { + return id[:i] + } + if len(id) > 8 { + return id[:8] + } + return id +} + +// formatStatusBattery renders a cached reading, or a dash when unknown. +func formatStatusBattery(b *ipc.StatusBattery) string { + if b == nil { + return "—" + } + if b.Charging { + return fmt.Sprintf("%d%%+", b.Charge) + } + return fmt.Sprintf("%d%%", b.Charge) +} + +// formatSeenAge renders an RFC3339 timestamp as a relative age. +func formatSeenAge(ts string) string { + if ts == "" { + return "never" + } + t, err := time.Parse(time.RFC3339, ts) + if err != nil { + return ts + } + d := time.Since(t) + switch { + case d < 0: + return "now" + case d < time.Minute: + return fmt.Sprintf("%ds ago", int(d.Seconds())) + case d < time.Hour: + return fmt.Sprintf("%dm ago", int(d.Minutes())) + case d < 24*time.Hour: + return fmt.Sprintf("%dh ago", int(d.Hours())) + default: + return t.Format("2006-01-02") + } +} + +func orDash(s string) string { + if s == "" { + return "—" + } + return s +} diff --git a/cmd/kcd/cli_status_test.go b/cmd/kcd/cli_status_test.go new file mode 100644 index 0000000..2e9e753 --- /dev/null +++ b/cmd/kcd/cli_status_test.go @@ -0,0 +1,61 @@ +package main + +import ( + "strings" + "testing" + "time" + + "github.com/bethropolis/kcd/internal/ipc" +) + +func TestFormatStatus(t *testing.T) { + st := &ipc.StatusResponse{ + Version: "v1.19.0", + UptimeHuman: "2h 14m", + SocketPath: "/run/user/1000/kcd/kcd.sock", + ConfigPath: "/home/user/.config/kcd/kcd.toml", + TCPPort: 1716, + Plugins: []string{"Pair", "Battery"}, + DeviceCount: 2, + ConnectedCount: 1, + Devices: []ipc.StatusDevice{ + { + ID: "9a5c23ea_7195_4da1_b766_282b7256a02d", Name: "BETHRÖ", + Type: "phone", State: "PAIRED", Connected: true, + Addr: "192.168.1.134:1716", + Battery: &ipc.StatusBattery{Charge: 78, Charging: true}, + LastSeen: time.Now().Add(-3 * time.Second).UTC().Format(time.RFC3339), + }, + { + ID: "deadbeef_0000_0000_0000_000000000000", Name: "Old Laptop", + Type: "laptop", State: "UNPAIRED", Connected: false, + }, + }, + } + out := formatStatus(st) + for _, want := range []string{ + "kcd v1.19.0 (up 2h 14m)", + "\nSocket: /run/user/1000/kcd/kcd.sock", + "Listen: tcp :1716", + "\nDevices: 2 known, 1 connected", + "BETHRÖ", "9a5c23ea", "PAIRED", "192.168.1.134:1716", "78%+", + "Old Laptop", "UNPAIRED", "—", "never", + "Plugins (2): Pair, Battery", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q:\n%s", want, out) + } + } +} + +func TestFormatSeenAge(t *testing.T) { + if got := formatSeenAge(""); got != "never" { + t.Errorf("empty = %q, want never", got) + } + if got := formatSeenAge("not-a-time"); got != "not-a-time" { + t.Errorf("garbage = %q, want passthrough", got) + } + if got := formatSeenAge(time.Now().Add(-90 * time.Second).UTC().Format(time.RFC3339)); got != "1m ago" { + t.Errorf("90s = %q, want 1m ago", got) + } +} diff --git a/cmd/kcd/main.go b/cmd/kcd/main.go index 5443d38..f99803e 100644 --- a/cmd/kcd/main.go +++ b/cmd/kcd/main.go @@ -7,7 +7,6 @@ import ( "log" "os" "os/signal" - "strings" "syscall" "time" @@ -195,11 +194,7 @@ func main() { fmt.Println(string(data)) return nil } - fmt.Printf("kcd %s — up %s\n", st.Version, st.UptimeHuman) - fmt.Printf("Socket: %s\n", st.SocketPath) - fmt.Printf("Config: %s\n", st.ConfigPath) - fmt.Printf("Devices: %d known, %d connected\n", st.DeviceCount, st.ConnectedCount) - fmt.Printf("Plugins: %s\n", strings.Join(st.Plugins, ", ")) + fmt.Print(formatStatus(st)) return nil }, }, diff --git a/docs/CLI.md b/docs/CLI.md index aa95d63..1db81cb 100644 --- a/docs/CLI.md +++ b/docs/CLI.md @@ -135,11 +135,18 @@ Show daemon runtime information. **Example output** - kcd v1.0.5 — up 3h 12m - Socket: /run/user/1000/kcd/kcd.sock - Config: /home/user/.config/kcd/kcd.toml - Devices: 2 known, 1 connected - Plugins: Battery, Clipboard, Notification, Share, ... + kcd v1.19.0 (up 2h 14m) + + Socket: /run/user/1000/kcd/kcd.sock + Config: /home/user/.config/kcd/kcd.toml + Listen: tcp :1716 + + Devices: 2 known, 1 connected + NAME ID TYPE STATE ADDR BATTERY LAST SEEN + BETHRÖ 9a5c23ea phone PAIRED 192.168.1.134:1716 78%+ 3s ago + Old Laptop deadbeef laptop UNPAIRED — — never + + Plugins (20): Pair, Battery, Clipboard, Notification, Share, ... --- diff --git a/internal/daemon/ipc_routes.go b/internal/daemon/ipc_routes.go index dbb54ae..5b6632e 100644 --- a/internal/daemon/ipc_routes.go +++ b/internal/daemon/ipc_routes.go @@ -115,11 +115,39 @@ func registerIPCRoutes(handler *ipc.Handler, cfg *config.Config, devices *device total := 0 connected := 0 + devInfos := make([]ipc.StatusDevice, 0) for _, d := range devices.List() { total++ - if d.IsConnected() { + isConnected := d.IsConnected() + if isConnected { connected++ } + info := ipc.StatusDevice{ + ID: d.ID(), + Name: d.Name(), + Type: d.Type, + State: d.State().String(), + Connected: isConnected, + } + if ip := d.RemoteIP(); ip != nil && isConnected { + info.Addr = ip.String() + if port := d.LastPort(); port > 0 { + info.Addr += fmt.Sprintf(":%d", port) + } + } else if ip := d.LastIP(); ip != nil { + info.Addr = ip.String() + } + if charge, charging := d.GetBattery(); d.HasBattery() { + info.Battery = &ipc.StatusBattery{ + Charge: charge, + Charging: charging, + AgeMs: d.BatteryAge().Milliseconds(), + } + } + if lastSeen := d.LastSeen(); !lastSeen.IsZero() { + info.LastSeen = lastSeen.UTC().Format(time.RFC3339) + } + devInfos = append(devInfos, info) } data, _ := json.Marshal(ipc.StatusResponse{ @@ -128,9 +156,11 @@ func registerIPCRoutes(handler *ipc.Handler, cfg *config.Config, devices *device UptimeHuman: uptimeHuman, SocketPath: cfg.SocketPath, ConfigPath: cfg.ConfigPath, + TCPPort: cfg.TCPPort, Plugins: pluginNames, DeviceCount: total, ConnectedCount: connected, + Devices: devInfos, }) return ipc.Response{OK: true, Data: data} }) diff --git a/internal/ipc/proto.go b/internal/ipc/proto.go index 4dbcd12..cbbe19b 100644 --- a/internal/ipc/proto.go +++ b/internal/ipc/proto.go @@ -128,16 +128,37 @@ type SMSAttachmentPayload struct { UniqueIdentifier string `json:"uniqueIdentifier"` } +// StatusBattery is a cached device battery reading in CmdStatus output. +type StatusBattery struct { + Charge int `json:"charge"` + Charging bool `json:"charging"` + AgeMs int64 `json:"ageMs"` +} + +// StatusDevice describes one known device in CmdStatus output. +type StatusDevice struct { + ID string `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + State string `json:"state"` + Connected bool `json:"connected"` + Addr string `json:"addr,omitempty"` + Battery *StatusBattery `json:"battery,omitempty"` + LastSeen string `json:"lastSeen,omitempty"` +} + // StatusResponse is returned by CmdStatus. type StatusResponse struct { - Version string `json:"version"` - StartedAt string `json:"startedAt"` - UptimeHuman string `json:"uptimeHuman"` - SocketPath string `json:"socketPath"` - ConfigPath string `json:"configPath"` - Plugins []string `json:"plugins"` - DeviceCount int `json:"deviceCount"` - ConnectedCount int `json:"connectedCount"` + Version string `json:"version"` + StartedAt string `json:"startedAt"` + UptimeHuman string `json:"uptimeHuman"` + SocketPath string `json:"socketPath"` + ConfigPath string `json:"configPath"` + TCPPort int `json:"tcpPort,omitempty"` + Plugins []string `json:"plugins"` + DeviceCount int `json:"deviceCount"` + ConnectedCount int `json:"connectedCount"` + Devices []StatusDevice `json:"devices,omitempty"` } // SftpInfoResponse carries cached SFTP connection details returned by CmdSftpInfo. From 81aedb0a6e45d9cb742cd6b225897492d27094e3 Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Fri, 18 Sep 2026 09:33:16 +0300 Subject: [PATCH 11/12] refactor: modularize oversized packages into focused files (#33) * refactor(battery): split plugin into types, handle and local files * refactor(systemvolume): split plugin into types, handle and backend files * refactor(mousepad): split plugin into types, dispatch and input files * refactor(pair): split plugin into types, handle and actions files * refactor(share): split plugin into types, handle and send files * refactor(sms): split plugin into types, handle, attachment and send files * refactor(notification): split plugin into types, handle, icon, desktop and actions files * refactor(clipboard): split plugin into types, backend, handle and push files * refactor(contacts): split plugin into types, request, store, sync and vcard files * refactor(sftp): split plugin into types, handle, request and mount files * refactor(mpris): split plugin into core, types, handle, local, remote, art and telephony files * refactor(device): split device into core, session, telemetry and discovery files * refactor(discovery): split into controller, broadcaster, listener and mdns files * refactor(ipc): extract watch stream handler into server_watch.go * refactor(config): split config into core, validate and store files * refactor(daemon): split transport into dial, handshake and reconnect files * refactor(client): split client by domain into focused files * refactor(cli): extract mpris parse and format helpers * refactor(cli): extract pair commands into cli_pair.go * chore(merge): port dev/next bugfixes into refactor split files Carry the P0 wire-protocol fixes across the file split so no behavior regresses: battery request/Request field, clipboard/sms partial cleanup, notification Dismiss, pair timestamp-seeded verification codes, share sidechannel options, SMS v2 schema with FlexBool read flags. --- cmd/kcd/cli_devices.go | 205 ---- cmd/kcd/cli_mpris.go | 53 - cmd/kcd/cli_mpris_format.go | 59 ++ cmd/kcd/cli_pair.go | 214 ++++ internal/config/config.go | 176 ---- internal/config/config_store.go | 130 +++ internal/config/config_validate.go | 60 ++ internal/daemon/transport.go | 306 ------ internal/daemon/transport_dial.go | 111 +++ internal/daemon/transport_handshake.go | 122 +++ internal/daemon/transport_reconnect.go | 120 +++ internal/device/device.go | 621 ------------ internal/device/device_core.go | 172 ++++ internal/device/device_discovery.go | 193 ++++ internal/device/device_session.go | 176 ++++ internal/device/device_telemetry.go | 108 ++ internal/discovery/broadcaster.go | 101 ++ internal/discovery/broadcaster_controller.go | 113 +++ internal/discovery/discovery.go | 380 -------- internal/discovery/listener.go | 93 ++ internal/discovery/listener_mdns.go | 105 ++ internal/ipc/server.go | 146 --- internal/ipc/server_watch.go | 153 +++ .../plugins/battery/{battery.go => handle.go} | 109 --- internal/plugins/battery/local.go | 63 ++ internal/plugins/battery/types.go | 57 ++ internal/plugins/clipboard/backend.go | 165 ++++ internal/plugins/clipboard/clipboard.go | 509 ---------- internal/plugins/clipboard/handle.go | 180 ++++ internal/plugins/clipboard/push.go | 122 +++ internal/plugins/clipboard/types.go | 77 ++ internal/plugins/contacts/contacts.go | 594 ----------- internal/plugins/contacts/request.go | 78 ++ internal/plugins/contacts/store.go | 122 +++ internal/plugins/contacts/sync.go | 231 +++++ internal/plugins/contacts/types.go | 114 +++ internal/plugins/contacts/vcard.go | 77 ++ internal/plugins/mousepad/dispatch.go | 63 ++ .../mousepad/{mousepad.go => input.go} | 159 --- internal/plugins/mousepad/types.go | 113 +++ internal/plugins/mpris/art.go | 165 ++++ internal/plugins/mpris/handle.go | 176 ++++ internal/plugins/mpris/local.go | 179 ++++ internal/plugins/mpris/mpris.go | 920 ------------------ internal/plugins/mpris/remote.go | 252 +++++ internal/plugins/mpris/telephony.go | 77 ++ internal/plugins/mpris/types.go | 123 +++ internal/plugins/notification/actions.go | 43 + internal/plugins/notification/desktop.go | 86 ++ internal/plugins/notification/handle.go | 139 +++ internal/plugins/notification/icon.go | 101 ++ internal/plugins/notification/notification.go | 463 --------- internal/plugins/notification/types.go | 127 +++ internal/plugins/pair/actions.go | 141 +++ internal/plugins/pair/handle.go | 142 +++ internal/plugins/pair/pair.go | 345 ------- internal/plugins/pair/types.go | 83 ++ internal/plugins/sftp/handle.go | 65 ++ internal/plugins/sftp/mount.go | 285 ++++++ internal/plugins/sftp/request.go | 205 ++++ internal/plugins/sftp/sftp.go | 601 ------------ internal/plugins/sftp/types.go | 76 ++ internal/plugins/share/handle.go | 147 +++ internal/plugins/share/send.go | 113 +++ internal/plugins/share/share.go | 338 ------- internal/plugins/share/types.go | 102 ++ internal/plugins/sms/attachment.go | 126 +++ internal/plugins/sms/handle.go | 95 ++ internal/plugins/sms/send.go | 75 ++ internal/plugins/sms/sms.go | 390 -------- internal/plugins/sms/types.go | 123 +++ internal/plugins/systemvolume/backend.go | 110 +++ internal/plugins/systemvolume/handle.go | 73 ++ internal/plugins/systemvolume/systemvolume.go | 232 ----- internal/plugins/systemvolume/types.go | 62 ++ pkg/client/client.go | 465 --------- pkg/client/client_actions.go | 97 ++ pkg/client/client_devices.go | 98 ++ pkg/client/client_mpris.go | 96 ++ pkg/client/client_sftp.go | 81 ++ pkg/client/client_sms_contacts.go | 62 ++ pkg/client/client_watch.go | 81 ++ 82 files changed, 7498 insertions(+), 7012 deletions(-) create mode 100644 cmd/kcd/cli_mpris_format.go create mode 100644 cmd/kcd/cli_pair.go create mode 100644 internal/config/config_store.go create mode 100644 internal/config/config_validate.go create mode 100644 internal/daemon/transport_dial.go create mode 100644 internal/daemon/transport_handshake.go create mode 100644 internal/daemon/transport_reconnect.go delete mode 100644 internal/device/device.go create mode 100644 internal/device/device_core.go create mode 100644 internal/device/device_discovery.go create mode 100644 internal/device/device_session.go create mode 100644 internal/device/device_telemetry.go create mode 100644 internal/discovery/broadcaster.go create mode 100644 internal/discovery/broadcaster_controller.go delete mode 100644 internal/discovery/discovery.go create mode 100644 internal/discovery/listener.go create mode 100644 internal/discovery/listener_mdns.go create mode 100644 internal/ipc/server_watch.go rename internal/plugins/battery/{battery.go => handle.go} (50%) create mode 100644 internal/plugins/battery/local.go create mode 100644 internal/plugins/battery/types.go create mode 100644 internal/plugins/clipboard/backend.go delete mode 100644 internal/plugins/clipboard/clipboard.go create mode 100644 internal/plugins/clipboard/handle.go create mode 100644 internal/plugins/clipboard/push.go create mode 100644 internal/plugins/clipboard/types.go delete mode 100644 internal/plugins/contacts/contacts.go create mode 100644 internal/plugins/contacts/request.go create mode 100644 internal/plugins/contacts/store.go create mode 100644 internal/plugins/contacts/sync.go create mode 100644 internal/plugins/contacts/types.go create mode 100644 internal/plugins/contacts/vcard.go create mode 100644 internal/plugins/mousepad/dispatch.go rename internal/plugins/mousepad/{mousepad.go => input.go} (53%) create mode 100644 internal/plugins/mousepad/types.go create mode 100644 internal/plugins/mpris/art.go create mode 100644 internal/plugins/mpris/handle.go create mode 100644 internal/plugins/mpris/local.go create mode 100644 internal/plugins/mpris/remote.go create mode 100644 internal/plugins/mpris/telephony.go create mode 100644 internal/plugins/mpris/types.go create mode 100644 internal/plugins/notification/actions.go create mode 100644 internal/plugins/notification/desktop.go create mode 100644 internal/plugins/notification/handle.go create mode 100644 internal/plugins/notification/icon.go delete mode 100644 internal/plugins/notification/notification.go create mode 100644 internal/plugins/notification/types.go create mode 100644 internal/plugins/pair/actions.go create mode 100644 internal/plugins/pair/handle.go delete mode 100644 internal/plugins/pair/pair.go create mode 100644 internal/plugins/pair/types.go create mode 100644 internal/plugins/sftp/handle.go create mode 100644 internal/plugins/sftp/mount.go create mode 100644 internal/plugins/sftp/request.go delete mode 100644 internal/plugins/sftp/sftp.go create mode 100644 internal/plugins/sftp/types.go create mode 100644 internal/plugins/share/handle.go create mode 100644 internal/plugins/share/send.go delete mode 100644 internal/plugins/share/share.go create mode 100644 internal/plugins/share/types.go create mode 100644 internal/plugins/sms/attachment.go create mode 100644 internal/plugins/sms/handle.go create mode 100644 internal/plugins/sms/send.go delete mode 100644 internal/plugins/sms/sms.go create mode 100644 internal/plugins/sms/types.go create mode 100644 internal/plugins/systemvolume/backend.go create mode 100644 internal/plugins/systemvolume/handle.go delete mode 100644 internal/plugins/systemvolume/systemvolume.go create mode 100644 internal/plugins/systemvolume/types.go create mode 100644 pkg/client/client_actions.go create mode 100644 pkg/client/client_devices.go create mode 100644 pkg/client/client_mpris.go create mode 100644 pkg/client/client_sftp.go create mode 100644 pkg/client/client_sms_contacts.go create mode 100644 pkg/client/client_watch.go diff --git a/cmd/kcd/cli_devices.go b/cmd/kcd/cli_devices.go index cf21d05..81bdbbd 100644 --- a/cmd/kcd/cli_devices.go +++ b/cmd/kcd/cli_devices.go @@ -3,14 +3,8 @@ package main import ( "encoding/json" "fmt" - "os" - "os/signal" - "strings" - "syscall" - "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/ipc" "github.com/bethropolis/kcd/internal/protocol" "github.com/urfave/cli/v2" ) @@ -68,205 +62,6 @@ var devicesCmd = &cli.Command{ }, } -var pairCmd = &cli.Command{ - - Name: "pair", - Usage: "Initiate pairing or accept incoming requests", - Description: `With a device ID: send a pair request to that device (or accept if they already requested). - -Without a device ID: enter listen mode to receive and verify incoming pairing requests.`, - ArgsUsage: "[device-id]", - Flags: []cli.Flag{ - &cli.BoolFlag{ - Name: "yes", - Aliases: []string{"y"}, - Usage: "Automatically accept incoming requests without confirmation (headless mode)", - }, - &cli.StringFlag{ - Name: "expected-fingerprint", - Usage: "Only accept a candidate whose TLS cert fingerprint matches (hex, colons optional)", - }, - &cli.BoolFlag{ - Name: "known-only", - Usage: "Only accept candidates already recorded in the known-devices file", - }, - }, - Action: func(c *cli.Context) error { - cl, err := getClient(c) - if err != nil { - return err - } - - if c.NArg() >= 1 { - targetID := c.Args().First() - if err := cl.Pair(targetID); err != nil { - return err - } - fmt.Printf("Pair request sent / accepted for %s\n", targetID) - return nil - } - - // Listen mode — wait for any incoming pair request - fmt.Println("Listening for pair requests… (Ctrl+C to cancel)") - - if c.Bool("yes") && c.String("expected-fingerprint") == "" && !c.Bool("known-only") { - fmt.Fprintln(os.Stderr, "WARNING: auto-accepting pairing requests from ANY device on the local network.") - fmt.Fprintln(os.Stderr, "Use --expected-fingerprint or --known-only to restrict which device may pair.") - } - - // Snapshot of previously seen devices for --known-only. A stranger - // that was never recorded in the state file is never auto-accepted. - var known map[string]bool - if c.Bool("known-only") { - known = loadKnownDeviceIDs() - } - - if err := cl.BroadcastStart(); err != nil { - return fmt.Errorf("failed to start broadcast: %w", err) - } - - // Stop broadcast on Ctrl+C or normal exit - sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) - defer func() { - signal.Stop(sigCh) - _ = cl.BroadcastStop() - }() - - type listenResult struct { - result *ipc.PairListenResult - err error - } - - // Keep waiting past the daemon's 60s pair_listen timeout so a slow - // phone-side accept doesn't force a restart. Ctrl+C cancels. - for { - resultCh := make(chan listenResult, 1) - go func() { - r, err := cl.PairListen() - resultCh <- listenResult{r, err} - }() - - select { - case <-sigCh: - fmt.Println("\nCancelled") - return nil - case r := <-resultCh: - if r.err != nil { - if strings.Contains(r.err.Error(), "timed out") { - fmt.Println("No pair requests yet, still listening… (Ctrl+C to cancel)") - continue - } - return r.err - } - - fmt.Printf("\nIncoming pair request from:\n") - fmt.Printf(" Device: %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) - if r.result.VerificationKey != "" { - fmt.Printf(" Verification code: %s\n", r.result.VerificationKey) - } - - // Headless / auto-accept flag - if c.Bool("yes") { - if err := checkPairCandidate(c.String("expected-fingerprint"), c.Bool("known-only"), r.result, known); err != nil { - fmt.Printf("Refusing candidate: %s\n", err) - _ = cl.Unpair(r.result.DeviceID) - fmt.Println("Still listening… (Ctrl+C to cancel)") - continue - } - if err := cl.Pair(r.result.DeviceID); err != nil { - return fmt.Errorf("failed to accept pairing: %w", err) - } - fmt.Printf("Paired with %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) - return nil - } - - // Interactive prompt (default: reject) - fmt.Print("\nAccept pairing? [y/N]: ") - var response string - fmt.Scanln(&response) - - response = strings.TrimSpace(strings.ToLower(response)) - if response == "y" || response == "yes" { - if err := checkPairCandidate(c.String("expected-fingerprint"), c.Bool("known-only"), r.result, known); err != nil { - fmt.Printf("Refusing candidate: %s\n", err) - _ = cl.Unpair(r.result.DeviceID) - return nil - } - if err := cl.Pair(r.result.DeviceID); err != nil { - return fmt.Errorf("failed to accept pairing: %w", err) - } - fmt.Printf("Paired with %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) - return nil - } - - // User rejected: reject and cancel request - _ = cl.Unpair(r.result.DeviceID) - fmt.Printf("Rejected pairing with %s\n", protocol.DisplayName(r.result.DeviceName)) - return nil - } - } - }, -} - -var unpairCmd = &cli.Command{ - - Name: "unpair", - Usage: "Revoke trust and unpair from a device", - ArgsUsage: "", - Action: func(c *cli.Context) error { - if c.NArg() < 1 { - return fmt.Errorf("missing device ID") - } - cl, err := getClient(c) - if err != nil { - return err - } - if err := cl.Unpair(c.Args().First()); err != nil { - return err - } - fmt.Println("Unpaired successfully") - return nil - }, -} - -// normalizeFingerprint strips separators and lowercases a hex fingerprint so -// user-supplied values (aa:bb:.., AA BB ..) compare against daemon hex. -func normalizeFingerprint(fp string) string { - fp = strings.ReplaceAll(fp, ":", "") - fp = strings.ReplaceAll(fp, " ", "") - return strings.ToLower(fp) -} - -// loadKnownDeviceIDs returns the set of device IDs recorded in the daemon's -// persisted state file. Missing or unreadable state means nothing is known. -func loadKnownDeviceIDs() map[string]bool { - known := make(map[string]bool) - infos, err := device.LoadDevices(config.StatePath()) - if err != nil { - return known - } - for _, info := range infos { - known[info.ID] = true - } - return known -} - -// checkPairCandidate enforces the --expected-fingerprint and --known-only -// constraints on a pairing candidate. Fail-closed: a candidate without a -// reported fingerprint never satisfies an expected fingerprint. -func checkPairCandidate(expectedFP string, knownOnly bool, result *ipc.PairListenResult, known map[string]bool) error { - if expectedFP != "" { - if got := normalizeFingerprint(result.Fingerprint); got == "" || got != normalizeFingerprint(expectedFP) { - return fmt.Errorf("candidate fingerprint does not match --expected-fingerprint") - } - } - if knownOnly && !known[result.DeviceID] { - return fmt.Errorf("candidate %s is not in the known-devices file (--known-only)", result.DeviceID) - } - return nil -} - func printDeviceTable(devices []device.DeviceInfo) { if len(devices) == 0 { fmt.Println("No devices found.") diff --git a/cmd/kcd/cli_mpris.go b/cmd/kcd/cli_mpris.go index 88d1377..e757cf1 100644 --- a/cmd/kcd/cli_mpris.go +++ b/cmd/kcd/cli_mpris.go @@ -5,8 +5,6 @@ import ( "fmt" "os" "strconv" - "strings" - "time" "github.com/urfave/cli/v2" ) @@ -268,54 +266,3 @@ var mprisCmd = &cli.Command{ }, }, } - -func shortID(id string) string { - if len(id) > 12 { - return id[:12] + "…" - } - return id -} - -func formatMs(ms int64) string { - if ms < 0 { - return "??:??" - } - totalSec := ms / 1000 - min := totalSec / 60 - sec := totalSec % 60 - return fmt.Sprintf("%d:%02d", min, sec) -} - -// parseSeek parses a seek offset string into milliseconds. -// Supported formats: +30s, -10s, 1m30s, 45 (bare seconds). -func parseSeek(s string) (int64, error) { - if strings.HasPrefix(s, "+") || strings.HasPrefix(s, "-") { - isNeg := strings.HasPrefix(s, "-") - rest := s[1:] - d, err := parseDuration(rest) - if err != nil { - return 0, fmt.Errorf("invalid offset %q", s) - } - if isNeg { - return -d, nil - } - return d, nil - } - d, err := parseDuration(s) - if err != nil { - return 0, fmt.Errorf("invalid offset %q", s) - } - return d, nil -} - -func parseDuration(s string) (int64, error) { - d, err := time.ParseDuration(s) - if err == nil { - return int64(d.Milliseconds()), nil - } - // Try bare seconds - if secs, err := strconv.ParseFloat(s, 64); err == nil { - return int64(secs * 1000), nil - } - return 0, fmt.Errorf("cannot parse %q as duration", s) -} diff --git a/cmd/kcd/cli_mpris_format.go b/cmd/kcd/cli_mpris_format.go new file mode 100644 index 0000000..93a8795 --- /dev/null +++ b/cmd/kcd/cli_mpris_format.go @@ -0,0 +1,59 @@ +package main + +import ( + "fmt" + "strconv" + "strings" + "time" +) + +func shortID(id string) string { + if len(id) > 12 { + return id[:12] + "…" + } + return id +} + +func formatMs(ms int64) string { + if ms < 0 { + return "??:??" + } + totalSec := ms / 1000 + min := totalSec / 60 + sec := totalSec % 60 + return fmt.Sprintf("%d:%02d", min, sec) +} + +// parseSeek parses a seek offset string into milliseconds. +// Supported formats: +30s, -10s, 1m30s, 45 (bare seconds). +func parseSeek(s string) (int64, error) { + if strings.HasPrefix(s, "+") || strings.HasPrefix(s, "-") { + isNeg := strings.HasPrefix(s, "-") + rest := s[1:] + d, err := parseDuration(rest) + if err != nil { + return 0, fmt.Errorf("invalid offset %q", s) + } + if isNeg { + return -d, nil + } + return d, nil + } + d, err := parseDuration(s) + if err != nil { + return 0, fmt.Errorf("invalid offset %q", s) + } + return d, nil +} + +func parseDuration(s string) (int64, error) { + d, err := time.ParseDuration(s) + if err == nil { + return int64(d.Milliseconds()), nil + } + // Try bare seconds + if secs, err := strconv.ParseFloat(s, 64); err == nil { + return int64(secs * 1000), nil + } + return 0, fmt.Errorf("cannot parse %q as duration", s) +} diff --git a/cmd/kcd/cli_pair.go b/cmd/kcd/cli_pair.go new file mode 100644 index 0000000..10783be --- /dev/null +++ b/cmd/kcd/cli_pair.go @@ -0,0 +1,214 @@ +package main + +import ( + "fmt" + "os" + "os/signal" + "strings" + "syscall" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/ipc" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/urfave/cli/v2" +) + +var pairCmd = &cli.Command{ + + Name: "pair", + Usage: "Initiate pairing or accept incoming requests", + Description: `With a device ID: send a pair request to that device (or accept if they already requested). + +Without a device ID: enter listen mode to receive and verify incoming pairing requests.`, + ArgsUsage: "[device-id]", + Flags: []cli.Flag{ + &cli.BoolFlag{ + Name: "yes", + Aliases: []string{"y"}, + Usage: "Automatically accept incoming requests without confirmation (headless mode)", + }, + &cli.StringFlag{ + Name: "expected-fingerprint", + Usage: "Only accept a candidate whose TLS cert fingerprint matches (hex, colons optional)", + }, + &cli.BoolFlag{ + Name: "known-only", + Usage: "Only accept candidates already recorded in the known-devices file", + }, + }, + Action: func(c *cli.Context) error { + cl, err := getClient(c) + if err != nil { + return err + } + + if c.NArg() >= 1 { + targetID := c.Args().First() + if err := cl.Pair(targetID); err != nil { + return err + } + fmt.Printf("Pair request sent / accepted for %s\n", targetID) + return nil + } + + // Listen mode — wait for any incoming pair request + fmt.Println("Listening for pair requests… (Ctrl+C to cancel)") + + if c.Bool("yes") && c.String("expected-fingerprint") == "" && !c.Bool("known-only") { + fmt.Fprintln(os.Stderr, "WARNING: auto-accepting pairing requests from ANY device on the local network.") + fmt.Fprintln(os.Stderr, "Use --expected-fingerprint or --known-only to restrict which device may pair.") + } + + // Snapshot of previously seen devices for --known-only. A stranger + // that was never recorded in the state file is never auto-accepted. + var known map[string]bool + if c.Bool("known-only") { + known = loadKnownDeviceIDs() + } + + if err := cl.BroadcastStart(); err != nil { + return fmt.Errorf("failed to start broadcast: %w", err) + } + + // Stop broadcast on Ctrl+C or normal exit + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + defer func() { + signal.Stop(sigCh) + _ = cl.BroadcastStop() + }() + + type listenResult struct { + result *ipc.PairListenResult + err error + } + + // Keep waiting past the daemon's 60s pair_listen timeout so a slow + // phone-side accept doesn't force a restart. Ctrl+C cancels. + for { + resultCh := make(chan listenResult, 1) + go func() { + r, err := cl.PairListen() + resultCh <- listenResult{r, err} + }() + + select { + case <-sigCh: + fmt.Println("\nCancelled") + return nil + case r := <-resultCh: + if r.err != nil { + if strings.Contains(r.err.Error(), "timed out") { + fmt.Println("No pair requests yet, still listening… (Ctrl+C to cancel)") + continue + } + return r.err + } + + fmt.Printf("\nIncoming pair request from:\n") + fmt.Printf(" Device: %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) + if r.result.VerificationKey != "" { + fmt.Printf(" Verification code: %s\n", r.result.VerificationKey) + } + + // Headless / auto-accept flag + if c.Bool("yes") { + if err := checkPairCandidate(c.String("expected-fingerprint"), c.Bool("known-only"), r.result, known); err != nil { + fmt.Printf("Refusing candidate: %s\n", err) + _ = cl.Unpair(r.result.DeviceID) + fmt.Println("Still listening… (Ctrl+C to cancel)") + continue + } + if err := cl.Pair(r.result.DeviceID); err != nil { + return fmt.Errorf("failed to accept pairing: %w", err) + } + fmt.Printf("Paired with %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) + return nil + } + + // Interactive prompt (default: reject) + fmt.Print("\nAccept pairing? [y/N]: ") + var response string + fmt.Scanln(&response) + + response = strings.TrimSpace(strings.ToLower(response)) + if response == "y" || response == "yes" { + if err := checkPairCandidate(c.String("expected-fingerprint"), c.Bool("known-only"), r.result, known); err != nil { + fmt.Printf("Refusing candidate: %s\n", err) + _ = cl.Unpair(r.result.DeviceID) + return nil + } + if err := cl.Pair(r.result.DeviceID); err != nil { + return fmt.Errorf("failed to accept pairing: %w", err) + } + fmt.Printf("Paired with %s (%s)\n", protocol.DisplayName(r.result.DeviceName), r.result.DeviceID) + return nil + } + + // User rejected: reject and cancel request + _ = cl.Unpair(r.result.DeviceID) + fmt.Printf("Rejected pairing with %s\n", protocol.DisplayName(r.result.DeviceName)) + return nil + } + } + }, +} + +var unpairCmd = &cli.Command{ + + Name: "unpair", + Usage: "Revoke trust and unpair from a device", + ArgsUsage: "", + Action: func(c *cli.Context) error { + if c.NArg() < 1 { + return fmt.Errorf("missing device ID") + } + cl, err := getClient(c) + if err != nil { + return err + } + if err := cl.Unpair(c.Args().First()); err != nil { + return err + } + fmt.Println("Unpaired successfully") + return nil + }, +} + +// normalizeFingerprint strips separators and lowercases a hex fingerprint so +// user-supplied values (aa:bb:.., AA BB ..) compare against daemon hex. +func normalizeFingerprint(fp string) string { + fp = strings.ReplaceAll(fp, ":", "") + fp = strings.ReplaceAll(fp, " ", "") + return strings.ToLower(fp) +} + +// loadKnownDeviceIDs returns the set of device IDs recorded in the daemon's +// persisted state file. Missing or unreadable state means nothing is known. +func loadKnownDeviceIDs() map[string]bool { + known := make(map[string]bool) + infos, err := device.LoadDevices(config.StatePath()) + if err != nil { + return known + } + for _, info := range infos { + known[info.ID] = true + } + return known +} + +// checkPairCandidate enforces the --expected-fingerprint and --known-only +// constraints on a pairing candidate. Fail-closed: a candidate without a +// reported fingerprint never satisfies an expected fingerprint. +func checkPairCandidate(expectedFP string, knownOnly bool, result *ipc.PairListenResult, known map[string]bool) error { + if expectedFP != "" { + if got := normalizeFingerprint(result.Fingerprint); got == "" || got != normalizeFingerprint(expectedFP) { + return fmt.Errorf("candidate fingerprint does not match --expected-fingerprint") + } + } + if knownOnly && !known[result.DeviceID] { + return fmt.Errorf("candidate %s is not in the known-devices file (--known-only)", result.DeviceID) + } + return nil +} diff --git a/internal/config/config.go b/internal/config/config.go index be0f80a..42212fb 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -3,12 +3,9 @@ package config import ( - "crypto/rand" "fmt" "os" "path/filepath" - "strings" - "time" "github.com/BurntSushi/toml" ) @@ -112,176 +109,3 @@ func Load(path string) (*Config, error) { } return cfg, nil } - -// Validate checks required fields and returns an error if any are invalid. -func (c *Config) Validate() error { - if err := c.validateDurations(); err != nil { - return err - } - if c.DeviceName == "" { - return fmt.Errorf("config: device_name is required") - } - switch c.DeviceType { - case "desktop", "laptop", "phone", "tablet", "tv": - // valid - default: - return fmt.Errorf("config: invalid device_type %q (expected desktop, laptop, phone, tablet, tv)", c.DeviceType) - } - if c.TCPPort < 1 || c.TCPPort > 65535 { - return fmt.Errorf("config: tcp_port must be 1-65535, got %d", c.TCPPort) - } - - // Battery urgency validation - for _, u := range []string{c.Battery.LowUrgency, c.Battery.FullUrgency} { - if u != "" { - switch u { - case "low", "normal", "critical": - default: - return fmt.Errorf("config: invalid battery urgency %q (expected low, normal, critical)", u) - } - } - } - - // Share port range validation - if c.Share.PortMin < 1 || c.Share.PortMin > 65535 || c.Share.PortMax < 1 || c.Share.PortMax > 65535 { - return fmt.Errorf("config: share ports must be 1-65535") - } - if c.Share.PortMin > c.Share.PortMax { - return fmt.Errorf("config: share port_min (%d) cannot be greater than port_max (%d)", c.Share.PortMin, c.Share.PortMax) - } - - // Prune threshold validation - if c.PruneStaleThreshold != "" { - if _, err := time.ParseDuration(c.PruneStaleThreshold); err != nil { - return fmt.Errorf("config: invalid prune_stale_threshold %q: %w", c.PruneStaleThreshold, err) - } - } - - // Mousepad backend validation - switch c.Mousepad.Backend { - case "auto", "ydotool", "xdotool", "uinput": - default: - return fmt.Errorf("config: invalid mousepad backend %q (expected auto, ydotool, xdotool)", c.Mousepad.Backend) - } - - return nil -} - -// EnsureDeviceID generates a UUIDv4-style device ID if one is not already set, -// and writes the updated config back to the given path. -func (c *Config) EnsureDeviceID(configPath string) error { - if c.DeviceID != "" { - return nil - } - - id, err := generateDeviceID() - if err != nil { - return fmt.Errorf("config: generate device id: %w", err) - } - c.DeviceID = id - - // Persist the generated ID if a config path is provided. - if configPath != "" { - if err := c.Save(configPath); err != nil { - return fmt.Errorf("config: save after generating device id: %w", err) - } - } - return nil -} - -// Save writes the config to a TOML file, creating parent directories as needed. -func (c *Config) Save(path string) error { - dir := filepath.Dir(path) - if err := os.MkdirAll(dir, 0700); err != nil { - return fmt.Errorf("config: create dir %s: %w", dir, err) - } - - f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0600) - if err != nil { - return fmt.Errorf("config: open %s: %w", path, err) - } - defer f.Close() - - enc := toml.NewEncoder(f) - if err := enc.Encode(c); err != nil { - return fmt.Errorf("config: encode: %w", err) - } - return nil -} - -// StatePath returns the path to the device state file. -func StatePath() string { - stateHome := os.Getenv("XDG_STATE_HOME") - if stateHome == "" { - home, _ := os.UserHomeDir() - stateHome = filepath.Join(home, ".local", "state") - } - return filepath.Join(stateHome, "kcd", "devices.json") -} - -// DefaultConfigPath returns the default config file path. -func DefaultConfigPath() string { - configHome := os.Getenv("XDG_CONFIG_HOME") - if configHome == "" { - home, _ := os.UserHomeDir() - configHome = filepath.Join(home, ".config") - } - return filepath.Join(configHome, "kcd", "kcd.toml") -} - -// DefaultSocketPath returns the default IPC socket path. -func DefaultSocketPath() string { - uid := fmt.Sprintf("%d", os.Getuid()) - rtDir := os.Getenv("XDG_RUNTIME_DIR") - if rtDir == "" { - rtDir = filepath.Join("/run/user", uid) - } - return filepath.Join(rtDir, "kcd", "kcd.sock") -} - -// configPath returns a path in the kcd config directory. -func configPath(filename string, isRuntime bool) string { - if isRuntime { - return filepath.Join(DefaultSocketPath()) - } - dir := filepath.Dir(DefaultConfigPath()) - return filepath.Join(dir, filename) -} - -// NotificationConfig controls notification branding and per-app filtering. -// The reserved app_name key is not a filter. -type NotificationConfig map[string]string - -// AppName returns the desktop notification application name. -func (c NotificationConfig) AppName() string { - if name := c["app_name"]; name != "" { - return name - } - return "KDE Connect" -} - -// Filters returns an independent copy containing only per-app filter entries. -func (c NotificationConfig) Filters() NotificationConfig { - filters := make(NotificationConfig, len(c)) - for app, action := range c { - if app != "app_name" { - filters[app] = action - } - } - return filters -} - -// generateDeviceID produces a UUIDv4 with dashes replaced by underscores. -func generateDeviceID() (string, error) { - var uuid [16]byte - if _, err := rand.Read(uuid[:]); err != nil { - return "", err - } - uuid[6] = (uuid[6] & 0x0f) | 0x40 - uuid[8] = (uuid[8] & 0x3f) | 0x80 - - s := fmt.Sprintf("%x-%x-%x-%x-%x", - uuid[0:4], uuid[4:6], uuid[6:8], uuid[8:10], uuid[10:16]) - - return strings.ReplaceAll(s, "-", "_"), nil -} diff --git a/internal/config/config_store.go b/internal/config/config_store.go new file mode 100644 index 0000000..b03ec45 --- /dev/null +++ b/internal/config/config_store.go @@ -0,0 +1,130 @@ +package config + +import ( + "crypto/rand" + "fmt" + "os" + "path/filepath" + "strings" + + "github.com/BurntSushi/toml" +) + +// EnsureDeviceID generates a UUIDv4-style device ID if one is not already set, +// and writes the updated config back to the given path. +func (c *Config) EnsureDeviceID(configPath string) error { + if c.DeviceID != "" { + return nil + } + + id, err := generateDeviceID() + if err != nil { + return fmt.Errorf("config: generate device id: %w", err) + } + c.DeviceID = id + + // Persist the generated ID if a config path is provided. + if configPath != "" { + if err := c.Save(configPath); err != nil { + return fmt.Errorf("config: save after generating device id: %w", err) + } + } + return nil +} + +// Save writes the config to a TOML file, creating parent directories as needed. +func (c *Config) Save(path string) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0700); err != nil { + return fmt.Errorf("config: create dir %s: %w", dir, err) + } + + f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0600) + if err != nil { + return fmt.Errorf("config: open %s: %w", path, err) + } + defer f.Close() + + enc := toml.NewEncoder(f) + if err := enc.Encode(c); err != nil { + return fmt.Errorf("config: encode: %w", err) + } + return nil +} + +// StatePath returns the path to the device state file. +func StatePath() string { + stateHome := os.Getenv("XDG_STATE_HOME") + if stateHome == "" { + home, _ := os.UserHomeDir() + stateHome = filepath.Join(home, ".local", "state") + } + return filepath.Join(stateHome, "kcd", "devices.json") +} + +// DefaultConfigPath returns the default config file path. +func DefaultConfigPath() string { + configHome := os.Getenv("XDG_CONFIG_HOME") + if configHome == "" { + home, _ := os.UserHomeDir() + configHome = filepath.Join(home, ".config") + } + return filepath.Join(configHome, "kcd", "kcd.toml") +} + +// DefaultSocketPath returns the default IPC socket path. +func DefaultSocketPath() string { + uid := fmt.Sprintf("%d", os.Getuid()) + rtDir := os.Getenv("XDG_RUNTIME_DIR") + if rtDir == "" { + rtDir = filepath.Join("/run/user", uid) + } + return filepath.Join(rtDir, "kcd", "kcd.sock") +} + +// configPath returns a path in the kcd config directory. +func configPath(filename string, isRuntime bool) string { + if isRuntime { + return filepath.Join(DefaultSocketPath()) + } + dir := filepath.Dir(DefaultConfigPath()) + return filepath.Join(dir, filename) +} + +// NotificationConfig controls notification branding and per-app filtering. +// The reserved app_name key is not a filter. +type NotificationConfig map[string]string + +// AppName returns the desktop notification application name. +func (c NotificationConfig) AppName() string { + if name := c["app_name"]; name != "" { + return name + } + return "KDE Connect" +} + +// Filters returns an independent copy containing only per-app filter entries. +func (c NotificationConfig) Filters() NotificationConfig { + filters := make(NotificationConfig, len(c)) + for app, action := range c { + if app != "app_name" { + filters[app] = action + } + } + return filters +} + +// generateDeviceID produces a UUIDv4 with dashes replaced by underscores. +func generateDeviceID() (string, error) { + var uuid [16]byte + if _, err := rand.Read(uuid[:]); err != nil { + return "", err + } + uuid[6] = (uuid[6] & 0x0f) | 0x40 + uuid[8] = (uuid[8] & 0x3f) | 0x80 + + s := fmt.Sprintf("%x-%x-%x-%x-%x", + uuid[0:4], uuid[4:6], uuid[6:8], uuid[8:10], uuid[10:16]) + + return strings.ReplaceAll(s, "-", "_"), nil +} diff --git a/internal/config/config_validate.go b/internal/config/config_validate.go new file mode 100644 index 0000000..c1992cf --- /dev/null +++ b/internal/config/config_validate.go @@ -0,0 +1,60 @@ +package config + +import ( + "fmt" + "time" +) + +// Validate checks required fields and returns an error if any are invalid. +func (c *Config) Validate() error { + if err := c.validateDurations(); err != nil { + return err + } + if c.DeviceName == "" { + return fmt.Errorf("config: device_name is required") + } + switch c.DeviceType { + case "desktop", "laptop", "phone", "tablet", "tv": + // valid + default: + return fmt.Errorf("config: invalid device_type %q (expected desktop, laptop, phone, tablet, tv)", c.DeviceType) + } + if c.TCPPort < 1 || c.TCPPort > 65535 { + return fmt.Errorf("config: tcp_port must be 1-65535, got %d", c.TCPPort) + } + + // Battery urgency validation + for _, u := range []string{c.Battery.LowUrgency, c.Battery.FullUrgency} { + if u != "" { + switch u { + case "low", "normal", "critical": + default: + return fmt.Errorf("config: invalid battery urgency %q (expected low, normal, critical)", u) + } + } + } + + // Share port range validation + if c.Share.PortMin < 1 || c.Share.PortMin > 65535 || c.Share.PortMax < 1 || c.Share.PortMax > 65535 { + return fmt.Errorf("config: share ports must be 1-65535") + } + if c.Share.PortMin > c.Share.PortMax { + return fmt.Errorf("config: share port_min (%d) cannot be greater than port_max (%d)", c.Share.PortMin, c.Share.PortMax) + } + + // Prune threshold validation + if c.PruneStaleThreshold != "" { + if _, err := time.ParseDuration(c.PruneStaleThreshold); err != nil { + return fmt.Errorf("config: invalid prune_stale_threshold %q: %w", c.PruneStaleThreshold, err) + } + } + + // Mousepad backend validation + switch c.Mousepad.Backend { + case "auto", "ydotool", "xdotool", "uinput": + default: + return fmt.Errorf("config: invalid mousepad backend %q (expected auto, ydotool, xdotool)", c.Mousepad.Backend) + } + + return nil +} diff --git a/internal/daemon/transport.go b/internal/daemon/transport.go index 16fbcd9..fa42f37 100644 --- a/internal/daemon/transport.go +++ b/internal/daemon/transport.go @@ -9,7 +9,6 @@ import ( "sync" "time" - "github.com/bethropolis/kcd/internal/cert" "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/discovery" @@ -19,100 +18,6 @@ import ( "go.uber.org/zap" ) -// validDialPort reports whether a discovery-advertised TCP port is usable. -// Port 0 and out-of-range values come from malformed or hostile packets and -// must never reach the dialer (port 0 would dial ":0"). -func validDialPort(port int) bool { - return port > 0 && port <= 65535 -} - -// DialDevice connects to a device at IP:port. Unless force is set, dials -// inside reconnectCooldown are skipped so post-roam bursts can't complete -// near-simultaneously and churn the peer's duplicate resolution. -// Explicit user actions (pair intent, manual connect) pass force=true. -func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID string, targetProto int, identity *protocol.Packet, cfg *tls.Config, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, force bool, opts *config.Config) { - if targetIP == nil || !validDialPort(targetPort) { - logger.Debug("refusing to dial invalid target", - zap.String("device_id", targetID), - zap.String("ip", targetIP.String()), - zap.Int("port", targetPort)) - return - } - if !force { - if dev, ok := devices.Get(targetID); ok && dev.InCooldown() { - logger.Debug("skipping dial inside reconnect cooldown", - zap.String("device_id", targetID), - zap.String("ip", targetIP.String())) - return - } - } - addr := fmt.Sprintf("%s:%d", targetIP, targetPort) - logger.Debug("dialing discovered device", zap.String("device_id", targetID), zap.String("addr", addr)) - - dialer := &net.Dialer{ - Timeout: config.Duration(opts.Network.DialTimeout), - } - conn, err := dialer.DialContext(ctx, "tcp", addr) - if err != nil { - logger.Debug("failed to dial peer", zap.Error(err)) - return - } - if tcpConn, ok := conn.(*net.TCPConn); ok { - if err := tcpConn.SetKeepAliveConfig(net.KeepAliveConfig{ - Enable: true, - Idle: 30 * time.Second, - Interval: 10 * time.Second, - Count: 3, - }); err != nil { - _ = tcpConn.SetKeepAlive(true) - _ = tcpConn.SetKeepAlivePeriod(30 * time.Second) - } - } - - var myID protocol.IdentityBody - json.Unmarshal(identity.Body, &myID) - // Clamp the echoed protocol version: discovery bodies are unauthenticated - // and a garbage value here would only confuse the peer. - if targetProto <= 0 || targetProto > protocol.ProtocolVersion { - targetProto = protocol.ProtocolVersion - } - // TargetDeviceID stays empty (and thus absent on the wire) when the - // target is unknown, e.g. an explicit connect-by-IP: stock peers - // close pre-TLS identities addressed to any other device ID. - preTlsId := protocol.IdentityBody{ - DeviceID: myID.DeviceID, - DeviceName: myID.DeviceName, - DeviceType: myID.DeviceType, - ProtocolVersion: myID.ProtocolVersion, - TCPPort: myID.TCPPort, - TargetDeviceID: targetID, - TargetProtocolVersion: targetProto, - } - preTlsPkt, _ := protocol.NewPacket(protocol.TypeIdentity, preTlsId) - - if err := transport.WritePlaintextPacket(conn, preTlsPkt); err != nil { - conn.Close() - return - } - - // KDE Connect inverts TLS roles: TCP client acts as TLS server - tlsConn := tls.Server(conn, cfg) - handshakeCtx, cancel := context.WithTimeout(ctx, config.Duration(opts.Network.HandshakeTimeout)) - defer cancel() - if err := tlsConn.HandshakeContext(handshakeCtx); err != nil { - tlsConn.Close() - logger.Debug("tls handshake failed", zap.Error(err)) - return - } - - transConn := transport.NewConn(tlsConn) - // Ensure the connection is closed if handleNewConnection fails mid-setup - if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger, opts); err != nil { - logger.Debug("new connection setup failed", zap.Error(err)) - transConn.Close() - } -} - // shouldEphemeralClose reports whether an active connection that exists // only for the discovery handshake should be closed on a fresh sighting. // Connections are kept when pairing mode is active, when the device is @@ -376,214 +281,3 @@ func runTransport(ctx context.Context, cfg *tls.Config, bc *discovery.Broadcaste <-ctx.Done() } - -// Returning an error ensures the caller can close the connection if it fails mid-setup. -func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, cfg *tls.Config, logger *zap.Logger, opts *config.Config) error { - if err := conn.WritePacket(identity); err != nil { - return fmt.Errorf("failed to send identity: %w", err) - } - - peerPkt, err := conn.ReadPacket() - if err != nil { - return fmt.Errorf("failed to read peer identity: %w", err) - } - defer protocol.ReleasePacket(peerPkt) - - if peerPkt.Type != protocol.TypeIdentity { - return fmt.Errorf("expected identity packet, got %s", peerPkt.Type) - } - - var peerBody protocol.IdentityBody - if err := json.Unmarshal(peerPkt.Body, &peerBody); err != nil { - return fmt.Errorf("failed to unmarshal peer identity: %w", err) - } - - peerCert := conn.PeerCert() - if peerCert == nil { - return fmt.Errorf("no peer certificate presented") - } - - certCN := peerCert.Subject.CommonName - if certCN != peerBody.DeviceID { - return fmt.Errorf("certificate CN (%s) doesn't match device ID (%s)", certCN, peerBody.DeviceID) - } - - dev, ok := devices.Get(peerBody.DeviceID) - safeDeviceName := protocol.SanitizeDeviceName(peerBody.DeviceName) - if !ok { - dev = device.NewDevice(peerBody.DeviceID, safeDeviceName, peerBody.DeviceType, logger) - devices.Add(dev) - } else { - dev.SetName(safeDeviceName) - } - dev.SetLastSeen(time.Now()) - - certFP := cert.Fingerprint(peerCert) - if dev.State() == device.StatePaired && dev.CertFP != "" { - if dev.CertFP != certFP { - return fmt.Errorf("certificate fingerprint mismatch (possible MITM)") - } - } else { - dev.CertFP = certFP - } - - // Remember the peer's listening port from the authenticated exchange so - // paired auto-dials use a known-good target instead of trusting future - // (unauthenticated) discovery announcements. - if validDialPort(peerBody.TCPPort) { - dev.SetLastPort(peerBody.TCPPort) - } - - dev.IncomingCaps = peerBody.IncomingCapabilities - dev.OutgoingCaps = peerBody.OutgoingCapabilities - - logger.Debug("device connected", - zap.String("device_id", peerBody.DeviceID), - zap.String("device_name", safeDeviceName), - zap.Int("protocol_version", peerBody.ProtocolVersion)) - - dispatch := func(ctx context.Context, sender *device.Device, pkt *protocol.Packet) bool { - return plugins.Dispatch(ctx, sender, pkt) - } - - onConnect := func(sender *device.Device) { - plugins.OnConnect(sender) - } - - onDisconnect := func(sender *device.Device) { - plugins.OnDisconnect(sender) - // Only attempt reconnection for paired devices whose last IP we know. - // Unpaired or manually-disconnected devices are left alone. - if sender.State() != device.StatePaired { - return - } - lastIP := sender.LastIP() - if lastIP == nil { - return - } - // A connection that lasted long enough was genuinely stable, so the - // next drop should start the backoff over. A flap (connection dies - // shortly after a successful dial, e.g. a dying peer) keeps the - // counter so the backoff keeps escalating instead of hammering at the - // 2s floor forever. - if sender.ConnectionAge() >= config.Duration(opts.Reconnect.FlapThreshold) { - sender.ResetReconnectAttempt() - } - // Prevent multiple concurrent reconnect goroutines for the same device. - if !sender.TryReconnect() { - logger.Debug("auto-reconnect: already reconnecting, skipping", - zap.String("device_id", sender.ID())) - return - } - go reconnectWithBackoff(ctx, sender, lastIP, identity, cfg, devices, plugins, localDeviceID, logger, opts) - } - - dev.Connect(ctx, conn, dispatch, onConnect, onDisconnect) - return nil -} - -// reconnectWithBackoff dials a paired device after it disconnects, using -// exponential backoff up to 5 minutes between attempts. It stops as soon as: -// - the device reconnects (IsConnected becomes true), or -// - the daemon context is cancelled, or -// - the device is unpaired. -// -// A fresh connection coming in from the phone side (inbound TCP) will set -// IsConnected, causing the loop to exit cleanly without a duplicate dial. -func reconnectWithBackoff( - ctx context.Context, - dev *device.Device, - ip net.IP, - identity *protocol.Packet, - cfg *tls.Config, - devices *device.Registry, - plugins *plugin.Registry, - localDeviceID string, - logger *zap.Logger, - opts *config.Config, -) { - maxBackoff := config.Duration(opts.Reconnect.MaxBackoff) - attempt := dev.ReconnectAttempt() - - defer dev.ReconnectDone() - - logger.Info("starting auto-reconnect", - zap.String("device_id", dev.ID()), - zap.String("device_name", dev.Name()), - zap.String("ip", ip.String()), - ) - - for { - // Stop if the daemon is shutting down. - if ctx.Err() != nil { - return - } - - // Stop if the device was unpaired while we were waiting. - if dev.State() != device.StatePaired { - logger.Debug("auto-reconnect: device no longer paired, stopping", - zap.String("device_id", dev.ID())) - return - } - - // Stop if the device already reconnected (inbound connection from phone). - if dev.IsConnected() { - logger.Debug("auto-reconnect: device already connected, stopping", - zap.String("device_id", dev.ID())) - return - } - - backoff := device.ReconnectBackoff(attempt, maxBackoff, config.Duration(opts.Reconnect.InitialBackoff)) - logger.Debug("auto-reconnect: waiting before next attempt", - zap.String("device_id", dev.ID()), - zap.Int("attempt", attempt+1), - zap.Duration("backoff", backoff), - ) - - select { - case <-ctx.Done(): - return - case <-time.After(backoff): - } - - // Re-check after the sleep — the phone may have connected inbound. - if dev.IsConnected() || dev.State() != device.StatePaired { - return - } - - logger.Info("auto-reconnect: dialling", - zap.String("device_id", dev.ID()), - zap.String("ip", ip.String()), - zap.Int("attempt", attempt+1), - ) - - // Prefer the peer's last advertised listening port over the - // default: the identity may carry a non-standard port (or none - // at all, in which case LastPort is 0 and we fall back). - port := reconnectPort(dev, opts.TCPPort) - DialDevice(ctx, ip, port, dev.ID(), protocol.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) - - if dev.IsConnected() { - logger.Info("auto-reconnect: succeeded", - zap.String("device_id", dev.ID()), - zap.Int("attempts", attempt+1), - ) - // Persist the counter before returning: if this connection flaps, - // the next reconnect cycle continues backing off rather than - // resetting to the 2s floor. onDisconnect resets it if the - // connection proved stable. - dev.SetReconnectAttempt(attempt + 1) - return - } - - attempt++ - } -} - -// reconnectPort prefers the authenticated peer port, falling back to local configuration. -func reconnectPort(dev *device.Device, fallback int) int { - if p := dev.LastPort(); validDialPort(p) { - return p - } - return fallback -} diff --git a/internal/daemon/transport_dial.go b/internal/daemon/transport_dial.go new file mode 100644 index 0000000..c7f38b8 --- /dev/null +++ b/internal/daemon/transport_dial.go @@ -0,0 +1,111 @@ +package daemon + +import ( + "context" + "crypto/tls" + "encoding/json" + "fmt" + "net" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// validDialPort reports whether a discovery-advertised TCP port is usable. +// Port 0 and out-of-range values come from malformed or hostile packets and +// must never reach the dialer (port 0 would dial ":0"). +func validDialPort(port int) bool { + return port > 0 && port <= 65535 +} + +// DialDevice connects to a device at IP:port. Unless force is set, dials +// inside reconnectCooldown are skipped so post-roam bursts can't complete +// near-simultaneously and churn the peer's duplicate resolution. +// Explicit user actions (pair intent, manual connect) pass force=true. +func DialDevice(ctx context.Context, targetIP net.IP, targetPort int, targetID string, targetProto int, identity *protocol.Packet, cfg *tls.Config, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, logger *zap.Logger, force bool, opts *config.Config) { + if targetIP == nil || !validDialPort(targetPort) { + logger.Debug("refusing to dial invalid target", + zap.String("device_id", targetID), + zap.String("ip", targetIP.String()), + zap.Int("port", targetPort)) + return + } + if !force { + if dev, ok := devices.Get(targetID); ok && dev.InCooldown() { + logger.Debug("skipping dial inside reconnect cooldown", + zap.String("device_id", targetID), + zap.String("ip", targetIP.String())) + return + } + } + addr := fmt.Sprintf("%s:%d", targetIP, targetPort) + logger.Debug("dialing discovered device", zap.String("device_id", targetID), zap.String("addr", addr)) + + dialer := &net.Dialer{ + Timeout: config.Duration(opts.Network.DialTimeout), + } + conn, err := dialer.DialContext(ctx, "tcp", addr) + if err != nil { + logger.Debug("failed to dial peer", zap.Error(err)) + return + } + if tcpConn, ok := conn.(*net.TCPConn); ok { + if err := tcpConn.SetKeepAliveConfig(net.KeepAliveConfig{ + Enable: true, + Idle: 30 * time.Second, + Interval: 10 * time.Second, + Count: 3, + }); err != nil { + _ = tcpConn.SetKeepAlive(true) + _ = tcpConn.SetKeepAlivePeriod(30 * time.Second) + } + } + + var myID protocol.IdentityBody + json.Unmarshal(identity.Body, &myID) + // Clamp the echoed protocol version: discovery bodies are unauthenticated + // and a garbage value here would only confuse the peer. + if targetProto <= 0 || targetProto > protocol.ProtocolVersion { + targetProto = protocol.ProtocolVersion + } + // TargetDeviceID stays empty (and thus absent on the wire) when the + // target is unknown, e.g. an explicit connect-by-IP: stock peers + // close pre-TLS identities addressed to any other device ID. + preTlsId := protocol.IdentityBody{ + DeviceID: myID.DeviceID, + DeviceName: myID.DeviceName, + DeviceType: myID.DeviceType, + ProtocolVersion: myID.ProtocolVersion, + TCPPort: myID.TCPPort, + TargetDeviceID: targetID, + TargetProtocolVersion: targetProto, + } + preTlsPkt, _ := protocol.NewPacket(protocol.TypeIdentity, preTlsId) + + if err := transport.WritePlaintextPacket(conn, preTlsPkt); err != nil { + conn.Close() + return + } + + // KDE Connect inverts TLS roles: TCP client acts as TLS server + tlsConn := tls.Server(conn, cfg) + handshakeCtx, cancel := context.WithTimeout(ctx, config.Duration(opts.Network.HandshakeTimeout)) + defer cancel() + if err := tlsConn.HandshakeContext(handshakeCtx); err != nil { + tlsConn.Close() + logger.Debug("tls handshake failed", zap.Error(err)) + return + } + + transConn := transport.NewConn(tlsConn) + // Ensure the connection is closed if handleNewConnection fails mid-setup + if err := handleNewConnection(ctx, transConn, identity, devices, plugins, localDeviceID, cfg, logger, opts); err != nil { + logger.Debug("new connection setup failed", zap.Error(err)) + transConn.Close() + } +} diff --git a/internal/daemon/transport_handshake.go b/internal/daemon/transport_handshake.go new file mode 100644 index 0000000..340ece1 --- /dev/null +++ b/internal/daemon/transport_handshake.go @@ -0,0 +1,122 @@ +package daemon + +import ( + "context" + "crypto/tls" + "encoding/json" + "fmt" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// Returning an error ensures the caller can close the connection if it fails mid-setup. +func handleNewConnection(ctx context.Context, conn *transport.Conn, identity *protocol.Packet, devices *device.Registry, plugins *plugin.Registry, localDeviceID string, cfg *tls.Config, logger *zap.Logger, opts *config.Config) error { + if err := conn.WritePacket(identity); err != nil { + return fmt.Errorf("failed to send identity: %w", err) + } + + peerPkt, err := conn.ReadPacket() + if err != nil { + return fmt.Errorf("failed to read peer identity: %w", err) + } + defer protocol.ReleasePacket(peerPkt) + + if peerPkt.Type != protocol.TypeIdentity { + return fmt.Errorf("expected identity packet, got %s", peerPkt.Type) + } + + var peerBody protocol.IdentityBody + if err := json.Unmarshal(peerPkt.Body, &peerBody); err != nil { + return fmt.Errorf("failed to unmarshal peer identity: %w", err) + } + + peerCert := conn.PeerCert() + if peerCert == nil { + return fmt.Errorf("no peer certificate presented") + } + + certCN := peerCert.Subject.CommonName + if certCN != peerBody.DeviceID { + return fmt.Errorf("certificate CN (%s) doesn't match device ID (%s)", certCN, peerBody.DeviceID) + } + + dev, ok := devices.Get(peerBody.DeviceID) + safeDeviceName := protocol.SanitizeDeviceName(peerBody.DeviceName) + if !ok { + dev = device.NewDevice(peerBody.DeviceID, safeDeviceName, peerBody.DeviceType, logger) + devices.Add(dev) + } else { + dev.SetName(safeDeviceName) + } + dev.SetLastSeen(time.Now()) + + certFP := cert.Fingerprint(peerCert) + if dev.State() == device.StatePaired && dev.CertFP != "" { + if dev.CertFP != certFP { + return fmt.Errorf("certificate fingerprint mismatch (possible MITM)") + } + } else { + dev.CertFP = certFP + } + + // Remember the peer's listening port from the authenticated exchange so + // paired auto-dials use a known-good target instead of trusting future + // (unauthenticated) discovery announcements. + if validDialPort(peerBody.TCPPort) { + dev.SetLastPort(peerBody.TCPPort) + } + + dev.IncomingCaps = peerBody.IncomingCapabilities + dev.OutgoingCaps = peerBody.OutgoingCapabilities + + logger.Debug("device connected", + zap.String("device_id", peerBody.DeviceID), + zap.String("device_name", safeDeviceName), + zap.Int("protocol_version", peerBody.ProtocolVersion)) + + dispatch := func(ctx context.Context, sender *device.Device, pkt *protocol.Packet) bool { + return plugins.Dispatch(ctx, sender, pkt) + } + + onConnect := func(sender *device.Device) { + plugins.OnConnect(sender) + } + + onDisconnect := func(sender *device.Device) { + plugins.OnDisconnect(sender) + // Only attempt reconnection for paired devices whose last IP we know. + // Unpaired or manually-disconnected devices are left alone. + if sender.State() != device.StatePaired { + return + } + lastIP := sender.LastIP() + if lastIP == nil { + return + } + // A connection that lasted long enough was genuinely stable, so the + // next drop should start the backoff over. A flap (connection dies + // shortly after a successful dial, e.g. a dying peer) keeps the + // counter so the backoff keeps escalating instead of hammering at the + // 2s floor forever. + if sender.ConnectionAge() >= config.Duration(opts.Reconnect.FlapThreshold) { + sender.ResetReconnectAttempt() + } + // Prevent multiple concurrent reconnect goroutines for the same device. + if !sender.TryReconnect() { + logger.Debug("auto-reconnect: already reconnecting, skipping", + zap.String("device_id", sender.ID())) + return + } + go reconnectWithBackoff(ctx, sender, lastIP, identity, cfg, devices, plugins, localDeviceID, logger, opts) + } + + dev.Connect(ctx, conn, dispatch, onConnect, onDisconnect) + return nil +} diff --git a/internal/daemon/transport_reconnect.go b/internal/daemon/transport_reconnect.go new file mode 100644 index 0000000..5ace804 --- /dev/null +++ b/internal/daemon/transport_reconnect.go @@ -0,0 +1,120 @@ +package daemon + +import ( + "context" + "crypto/tls" + "net" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// reconnectWithBackoff dials a paired device after it disconnects, using +// exponential backoff up to 5 minutes between attempts. It stops as soon as: +// - the device reconnects (IsConnected becomes true), or +// - the daemon context is cancelled, or +// - the device is unpaired. +// +// A fresh connection coming in from the phone side (inbound TCP) will set +// IsConnected, causing the loop to exit cleanly without a duplicate dial. +func reconnectWithBackoff( + ctx context.Context, + dev *device.Device, + ip net.IP, + identity *protocol.Packet, + cfg *tls.Config, + devices *device.Registry, + plugins *plugin.Registry, + localDeviceID string, + logger *zap.Logger, + opts *config.Config, +) { + maxBackoff := config.Duration(opts.Reconnect.MaxBackoff) + attempt := dev.ReconnectAttempt() + + defer dev.ReconnectDone() + + logger.Info("starting auto-reconnect", + zap.String("device_id", dev.ID()), + zap.String("device_name", dev.Name()), + zap.String("ip", ip.String()), + ) + + for { + // Stop if the daemon is shutting down. + if ctx.Err() != nil { + return + } + + // Stop if the device was unpaired while we were waiting. + if dev.State() != device.StatePaired { + logger.Debug("auto-reconnect: device no longer paired, stopping", + zap.String("device_id", dev.ID())) + return + } + + // Stop if the device already reconnected (inbound connection from phone). + if dev.IsConnected() { + logger.Debug("auto-reconnect: device already connected, stopping", + zap.String("device_id", dev.ID())) + return + } + + backoff := device.ReconnectBackoff(attempt, maxBackoff, config.Duration(opts.Reconnect.InitialBackoff)) + logger.Debug("auto-reconnect: waiting before next attempt", + zap.String("device_id", dev.ID()), + zap.Int("attempt", attempt+1), + zap.Duration("backoff", backoff), + ) + + select { + case <-ctx.Done(): + return + case <-time.After(backoff): + } + + // Re-check after the sleep — the phone may have connected inbound. + if dev.IsConnected() || dev.State() != device.StatePaired { + return + } + + logger.Info("auto-reconnect: dialling", + zap.String("device_id", dev.ID()), + zap.String("ip", ip.String()), + zap.Int("attempt", attempt+1), + ) + + // Prefer the peer's last advertised listening port over the + // default: the identity may carry a non-standard port (or none + // at all, in which case LastPort is 0 and we fall back). + port := reconnectPort(dev, opts.TCPPort) + DialDevice(ctx, ip, port, dev.ID(), protocol.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false, opts) + + if dev.IsConnected() { + logger.Info("auto-reconnect: succeeded", + zap.String("device_id", dev.ID()), + zap.Int("attempts", attempt+1), + ) + // Persist the counter before returning: if this connection flaps, + // the next reconnect cycle continues backing off rather than + // resetting to the 2s floor. onDisconnect resets it if the + // connection proved stable. + dev.SetReconnectAttempt(attempt + 1) + return + } + + attempt++ + } +} + +// reconnectPort prefers the authenticated peer port, falling back to local configuration. +func reconnectPort(dev *device.Device, fallback int) int { + if p := dev.LastPort(); validDialPort(p) { + return p + } + return fallback +} diff --git a/internal/device/device.go b/internal/device/device.go deleted file mode 100644 index fbdbc6d..0000000 --- a/internal/device/device.go +++ /dev/null @@ -1,621 +0,0 @@ -package device - -import ( - "context" - "crypto/x509" - "net" - "sync" - "sync/atomic" - "time" - - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "github.com/bethropolis/kcd/internal/transport" - "go.uber.org/zap" -) - -// Device represents an active KDE Connect remote device. -type Device struct { - id string - name string - Type string - - IncomingCaps []string - OutgoingCaps []string - - state PairingState - CertFP string - - lastSeen time.Time - lastIP net.IP // cached from last successful connection; survives Disconnect - // lastPort is the tcpPort the peer last advertised over the authenticated - // (post-TLS) identity exchange. Used with lastIP as the dial target for - // paired devices so unauthenticated discovery packets can never redirect - // a paired auto-dial. Zero means unknown (fall back to 1716). - lastPort int - - // discoveryIP/discoveryPort remember where a device was last seen - // announcing itself (UDP/mDNS), even if we never opened a TCP - // connection to it. Used to dial on explicit user request - // (e.g. `kcd pair `) without auto-dialling strangers. - discoveryIP net.IP - discoveryPort int - - // pairDialRequested is the one-shot outbound dial trigger for an explicit - // `kcd pair ` request. It is consumed by the first discovery - // announcement after the request so the pair request can be delivered. - pairDialRequested atomic.Bool - - // pairIntentUntil is the Unix-nano deadline until which an explicit pair - // intent keeps a connection alive. Unlike pairDialRequested (consumed on - // first sighting), the intent survives dial/connect cycles until pairing - // starts, is rejected, succeeds, or the deadline (pairDialIntentTTL) - // expires — so a slow phone-side accept can't downgrade into an - // ephemeral-close flap. - pairIntentUntil atomic.Int64 - - // lastDiscoveryDial is when onDeviceFound last spawned a dial for this - // device. It throttles sighting-triggered redials so announcements - // (or a spoofed broadcast storm) can't cause a dial per packet. - lastDiscoveryDial time.Time - - // ephemeralDialed marks that this device already received its one - // ephemeral discovery dial for the current unpaired era. Ephemeral - // dials let a stranger complete the TCP identity exchange (so both - // sides list each other) without staying connected: the next sighting - // closes the socket again while the device is still unpaired. The - // marker is cleared when the device is explicitly unpaired/rejected, - // making it eligible again. Paired devices and pairing mode bypass it. - ephemeralDialed bool - - conn *transport.Conn - sendChan chan *protocol.Packet // buffered 32 - done chan struct{} - closeOnce sync.Once - - // lastConnect marks the last completed handshake; new handshakes - // inside reconnectCooldown are refused to starve duplicate bursts. - lastConnect time.Time - // lastSightedIP remembers the previous discovery sighting so only - // confirmed roams reset the reconnect backoff (see NoteSighting). - lastSightedIP net.IP - BatteryCharge int - IsCharging bool - - // batterySeen marks that at least one kdeconnect.battery packet was - // received. Until then the zero values above are not measurements — - // they must not be published (a fresh pair would otherwise report a - // stable, bogus 0% that no later packet corrects at steady charge). - batterySeen bool - // lastBatteryAt is when the last battery packet arrived, so clients - // can apply their own staleness rules (mirrors mediaAgeMs). - lastBatteryAt time.Time - - mu sync.RWMutex - - // reconnecting is an atomic flag preventing multiple concurrent - // auto-reconnect goroutines for this device. - reconnecting atomic.Bool - - // reconnectAttempt persists the auto-reconnect backoff counter across - // disconnect cycles. A connection that flaps (drops shortly after a - // successful dial) keeps the counter so the backoff escalates instead of - // resetting to the 2s floor; a stable connection resets it on drop. - reconnectAttempt int - - // connectStarted is when the most recent connection was established, - // used to detect flaps (connections that die too quickly to count as - // genuinely stable). - connectStarted time.Time - - // pluginDispatch routes incoming packets to registered plugins - pluginDispatch func(ctx context.Context, dev *Device, pkt *protocol.Packet) bool - onConnect func(dev *Device) - onDisconnect func(dev *Device) - - logger *zap.Logger - bus *events.Bus -} - -// NewDevice creates a new disconnected device instance. -// New devices start as Unpaired (not Unknown) so listings are unambiguous. -func NewDevice(id, name, dtype string, logger *zap.Logger) *Device { - return &Device{ - id: id, - name: name, - Type: dtype, - state: StateUnpaired, - sendChan: make(chan *protocol.Packet, 32), - done: make(chan struct{}), - logger: logger.With(zap.String("device_id", id)), - } -} - -// SetBus sets the event bus for the device. -func (d *Device) SetBus(bus *events.Bus) { - d.mu.Lock() - defer d.mu.Unlock() - d.bus = bus -} - -// reconnectCooldown refuses new handshakes this long after a completed -// one, so post-roam bursts can't complete near-simultaneously and churn -// the peer's duplicate resolution. Reference stacks rate-limit the same -// way (desktop 500ms, Android MILLIS_DELAY_BETWEEN_CONNECTIONS_TO_SAME_DEVICE -// 1000ms); 1s matches upstream Android. -const reconnectCooldown = 1 * time.Second - -// Connect establishes a connection for the device and starts the reader and writer loops. -// A new authenticated connection immediately replaces any existing one -// (matching LanDeviceLink::reset / LanLink.reset). The old socket is -// closed; its readLoop will exit and disconnectConn will ignore it -// because d.conn no longer points at it. -func (d *Device) Connect(ctx context.Context, conn *transport.Conn, dispatch func(context.Context, *Device, *protocol.Packet) bool, onConnect func(*Device), onDisconnect func(*Device)) { - d.mu.Lock() - d.pluginDispatch = dispatch - d.onConnect = onConnect - d.onDisconnect = onDisconnect - - var oldConn *transport.Conn - var oldAddr, newAddr string - if d.conn != nil { - oldConn = d.conn - oldDone := d.done - oldAddr = oldConn.RemoteAddr().String() - newAddr = conn.RemoteAddr().String() - // Close the old done so its writerLoop exits; the new - // writerLoop will own the fresh channel. - d.closeOnce.Do(func() { - if oldDone != nil { - close(oldDone) - } - }) - } - - d.conn = conn - - // Renew the send channel on (re)connect in case it was closed during disconnect. - d.sendChan = make(chan *protocol.Packet, 32) - d.done = make(chan struct{}) - d.closeOnce = sync.Once{} - bus := d.bus - d.mu.Unlock() - - if oldConn != nil { - _ = oldConn.Close() - d.logger.Debug("replacing superseded connection", - zap.String("old_addr", oldAddr), - zap.String("new_addr", newAddr)) - } - - d.logger.Info("device connected", zap.String("remote_addr", conn.RemoteAddr().String())) - if bus != nil { - bus.Publish(events.TypeDeviceConnected, d.id, map[string]interface{}{ - "name": d.name, - "type": d.Type, - }) - } - - if d.onConnect != nil { - d.onConnect(d) - } - - // Cache the remote IP before the loops start so it is available - // after Disconnect() sets d.conn to nil (used by auto-reconnect). - if tcpAddr, ok := conn.RemoteAddr().(*net.TCPAddr); ok { - d.mu.Lock() - d.lastIP = tcpAddr.IP - d.mu.Unlock() - } - - // connectStarted distinguishes quick drops from stable connections, - // and doubles as the duplicate-cooldown clock (see InCooldown). - d.mu.Lock() - d.connectStarted = time.Now() - d.lastConnect = d.connectStarted - d.mu.Unlock() - - go d.readLoop(ctx, conn) - go d.writerLoop(ctx) -} - -// Disconnect terminates the session and stops the loops. -func (d *Device) Disconnect() { - d.mu.Lock() - d.lastSeen = time.Now() - - if d.conn == nil { - d.mu.Unlock() - return - } - - d.logger.Info("device disconnected") - _ = d.conn.Close() - d.conn = nil - d.lastConnect = time.Time{} - - // Capture these to call outside the lock to prevent deadlocks! - onDisc := d.onDisconnect - bus := d.bus - - d.closeOnce.Do(func() { - if d.done != nil { - close(d.done) - } - }) - d.mu.Unlock() - - // External calls must happen outside the mutex - if onDisc != nil { - onDisc(d) - } - if bus != nil { - bus.Publish(events.TypeDeviceDisconnected, d.id, nil) - } -} - -// IsConnected returns whether the device currently has an active connection. -func (d *Device) IsConnected() bool { - d.mu.RLock() - defer d.mu.RUnlock() - return d.conn != nil -} - -// disconnectConn handles a session's death. If the conn that died is no -// longer the current d.conn (it was superseded by a newer authenticated -// connection via Connect), the event is ignored silently — matching -// LanDeviceLink::reset's `if (m_socket == socket)` guard. Only the -// current preferred session's death triggers the full disconnect. -func (d *Device) disconnectConn(conn *transport.Conn) { - d.mu.Lock() - - if d.conn != conn { - d.mu.Unlock() - d.logger.Debug("ignoring disconnect from superseded connection") - return - } - - d.lastSeen = time.Now() - d.logger.Info("device disconnected") - _ = d.conn.Close() - d.conn = nil - d.lastConnect = time.Time{} - - // Capture callbacks to execute outside the lock - onDisc := d.onDisconnect - bus := d.bus - - d.closeOnce.Do(func() { - if d.done != nil { - close(d.done) - } - }) - d.mu.Unlock() - - // External calls must happen outside the mutex to prevent deadlocks - if onDisc != nil { - onDisc(d) - } - if bus != nil { - bus.Publish(events.TypeDeviceDisconnected, d.id, nil) - } -} - -// UpdateBattery updates the battery state of the device. -func (d *Device) UpdateBattery(charge int, charging bool) { - d.mu.Lock() - d.BatteryCharge = charge - d.IsCharging = charging - d.batterySeen = true - d.lastBatteryAt = time.Now() - bus := d.bus - id := d.id - d.mu.Unlock() - - if bus != nil { - bus.Publish(events.TypeBatteryUpdate, id, map[string]interface{}{ - "charge": charge, - "charging": charging, - }) - } -} - -// GetBattery returns the current battery state of the device. -func (d *Device) GetBattery() (int, bool) { - d.mu.RLock() - defer d.mu.RUnlock() - return d.BatteryCharge, d.IsCharging -} - -// HasBattery reports whether at least one battery packet was received. -// Until then the charge values are zero-value defaults, not measurements. -func (d *Device) HasBattery() bool { - d.mu.RLock() - defer d.mu.RUnlock() - return d.batterySeen -} - -// BatteryAge returns how long ago the last battery packet arrived, or a -// negative duration when no packet was ever received. -func (d *Device) BatteryAge() time.Duration { - d.mu.RLock() - defer d.mu.RUnlock() - if !d.batterySeen { - return -1 - } - return time.Since(d.lastBatteryAt) -} - -// HasCapability checks if the device has a particular capability (incoming or outgoing). -func (d *Device) HasCapability(cap string) bool { - d.mu.RLock() - defer d.mu.RUnlock() - - for _, c := range d.IncomingCaps { - if c == cap { - return true - } - } - for _, c := range d.OutgoingCaps { - if c == cap { - return true - } - } - return false -} - -// RemoteIP returns the IP address of the connected peer, if available. -func (d *Device) RemoteIP() net.IP { - d.mu.RLock() - defer d.mu.RUnlock() - - if d.conn == nil { - return nil - } - addr := d.conn.RemoteAddr() - if tcpAddr, ok := addr.(*net.TCPAddr); ok { - return tcpAddr.IP - } - return nil -} - -// LastIP returns the IP address from the most recent successful connection. -// Unlike RemoteIP, this persists after the connection drops — safe to use -// from an OnDisconnect callback for auto-reconnect dialling. -func (d *Device) LastIP() net.IP { - d.mu.RLock() - defer d.mu.RUnlock() - return d.lastIP -} - -// PeerCert returns the validated certificate presented by the remote device. -// Returns nil if not connected or no certificate was presented. -func (d *Device) PeerCert() *x509.Certificate { - d.mu.RLock() - defer d.mu.RUnlock() - - if d.conn == nil { - return nil - } - return d.conn.PeerCert() -} -func (d *Device) ID() string { - return d.id -} -func (d *Device) Name() string { - d.mu.RLock() - defer d.mu.RUnlock() - return d.name -} -func (d *Device) SetName(n string) { - d.mu.Lock() - defer d.mu.Unlock() - d.name = n -} -func (d *Device) State() PairingState { - d.mu.RLock() - defer d.mu.RUnlock() - return d.state -} -func (d *Device) SetState(s PairingState) { - d.mu.Lock() - defer d.mu.Unlock() - d.state = s -} -func (d *Device) LastSeen() time.Time { - d.mu.RLock() - defer d.mu.RUnlock() - return d.lastSeen -} -func (d *Device) SetLastSeen(t time.Time) { - d.mu.Lock() - defer d.mu.Unlock() - d.lastSeen = t -} - -// SetDiscoveryAddr records where the device was last seen announcing itself. -func (d *Device) SetDiscoveryAddr(ip net.IP, port int) { - d.mu.Lock() - defer d.mu.Unlock() - d.discoveryIP = ip - d.discoveryPort = port -} - -// DiscoveryAddr returns the last-seen announcement address, or nil if unknown. -func (d *Device) DiscoveryAddr() (net.IP, int) { - d.mu.RLock() - defer d.mu.RUnlock() - return d.discoveryIP, d.discoveryPort -} - -// pairDialIntentTTL bounds how long an explicit `kcd pair ` intent pins -// a connection while pairing hasn't started yet. -const pairDialIntentTTL = 5 * time.Minute - -// RequestPairDial marks the device for a one-shot outbound dial on its next -// discovery announcement and arms the keep-alive intent until pairing starts, -// is rejected, succeeds, or the TTL expires. Used when the user explicitly -// runs `kcd pair ` for a device with no active connection. -func (d *Device) RequestPairDial(intentTTL ...time.Duration) { - ttl := pairDialIntentTTL - if len(intentTTL) > 0 && intentTTL[0] > 0 { - ttl = intentTTL[0] - } - d.pairDialRequested.Store(true) - d.pairIntentUntil.Store(time.Now().Add(ttl).UnixNano()) -} - -// ConsumePairDial reports and clears a pending explicit pair-dial request. -func (d *Device) ConsumePairDial() bool { - return d.pairDialRequested.CompareAndSwap(true, false) -} - -// PairDialPending reports whether an explicit pair-dial was requested, -// without clearing it. -func (d *Device) PairDialPending() bool { - return d.pairDialRequested.Load() -} - -// PairDialActive reports whether an explicit pair intent is still keeping -// the connection alive: requested and neither cleared nor expired. Unlike -// PairDialPending (the one-shot dial trigger, consumed on first sighting), -// this survives dial/connect cycles until pairing starts or ends. -func (d *Device) PairDialActive() bool { - return d.pairIntentUntil.Load() > time.Now().UnixNano() -} - -// ClearPairDial drops both the one-shot trigger and the keep-alive intent. -// Call it when pairing completes, is rejected/cancelled, or is unpaired. -func (d *Device) ClearPairDial() { - d.pairDialRequested.Store(false) - d.pairIntentUntil.Store(0) -} - -// LastPort returns the last authenticated tcpPort advertised by the peer, -// or 0 if unknown. -func (d *Device) LastPort() int { - d.mu.RLock() - defer d.mu.RUnlock() - return d.lastPort -} - -// SetLastPort records the peer's advertised listening port after a -// successful authenticated exchange. Callers must pass a validated port. -func (d *Device) SetLastPort(port int) { - d.mu.Lock() - defer d.mu.Unlock() - d.lastPort = port -} - -// SetLastIP records a dial target, used when restoring persisted state. -// A nil IP clears the target. -func (d *Device) SetLastIP(ip net.IP) { - d.mu.Lock() - defer d.mu.Unlock() - d.lastIP = ip -} - -// ShouldDiscoveryDial reports whether enough time has passed since the last -// discovery-triggered dial for this device, and marks this dial if so. It -// bounds redial storms to a stale LastIP (DHCP roam) or spoofed sightings. -func (d *Device) ShouldDiscoveryDial(minInterval time.Duration) bool { - d.mu.Lock() - defer d.mu.Unlock() - if time.Since(d.lastDiscoveryDial) < minInterval { - return false - } - d.lastDiscoveryDial = time.Now() - return true -} - -// MarkEphemeralDialed records that the one ephemeral discovery dial for the -// current unpaired era has been used. -func (d *Device) MarkEphemeralDialed() { - d.mu.Lock() - defer d.mu.Unlock() - d.ephemeralDialed = true -} - -// EphemeralDialed reports whether the ephemeral discovery dial was used. -func (d *Device) EphemeralDialed() bool { - d.mu.RLock() - defer d.mu.RUnlock() - return d.ephemeralDialed -} - -// ClearEphemeral makes the device eligible for a fresh ephemeral discovery -// dial (e.g. after an explicit unpair or rejection). -func (d *Device) ClearEphemeral() { - d.mu.Lock() - defer d.mu.Unlock() - d.ephemeralDialed = false -} - -// InCooldown reports whether a handshake completed too recently to start -// another one for this device. -func (d *Device) InCooldown() bool { - d.mu.RLock() - defer d.mu.RUnlock() - return !d.lastConnect.IsZero() && time.Since(d.lastConnect) < reconnectCooldown -} - -// NoteSighting records a discovery sighting while disconnected, reporting -// whether it confirms a genuine roam: same new address twice running. -// Single sightings prove nothing (IPv4/IPv6 alternation, AP flicker). -func (d *Device) NoteSighting(sighted net.IP) (roamed bool) { - d.mu.Lock() - defer d.mu.Unlock() - defer func() { d.lastSightedIP = sighted }() - if lastIP := d.lastIP; lastIP == nil || !lastIP.Equal(sighted) { - return sighted.Equal(d.lastSightedIP) - } - return false -} - -// TryReconnect attempts to mark the device as reconnecting. -// Returns true if this goroutine should proceed; false if another -// reconnect goroutine is already running. -func (d *Device) TryReconnect() bool { - return d.reconnecting.CompareAndSwap(false, true) -} - -// ReconnectDone marks the device as no longer reconnecting. -// Must be called (typically via defer) after a reconnect goroutine exits. -func (d *Device) ReconnectDone() { - d.reconnecting.Store(false) -} - -// ReconnectAttempt returns the persisted auto-reconnect backoff counter. -func (d *Device) ReconnectAttempt() int { - d.mu.RLock() - defer d.mu.RUnlock() - return d.reconnectAttempt -} - -// SetReconnectAttempt stores the auto-reconnect backoff counter so the next -// reconnect cycle (spawned when this connection drops) continues backing off -// instead of resetting to the initial floor. -func (d *Device) SetReconnectAttempt(n int) { - d.mu.Lock() - d.reconnectAttempt = n - d.mu.Unlock() -} - -// ResetReconnectAttempt clears the backoff counter after a stable connection -// (one that stayed up long enough to count as genuinely healthy) drops. -func (d *Device) ResetReconnectAttempt() { - d.mu.Lock() - d.reconnectAttempt = 0 - d.mu.Unlock() -} - -// ConnectionAge returns how long the current connection has been established, -// or 0 if the device has never connected in this process. -func (d *Device) ConnectionAge() time.Duration { - d.mu.RLock() - defer d.mu.RUnlock() - if d.connectStarted.IsZero() { - return 0 - } - return time.Since(d.connectStarted) -} diff --git a/internal/device/device_core.go b/internal/device/device_core.go new file mode 100644 index 0000000..603e102 --- /dev/null +++ b/internal/device/device_core.go @@ -0,0 +1,172 @@ +package device + +import ( + "context" + "net" + "sync" + "sync/atomic" + "time" + + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// Device represents an active KDE Connect remote device. +type Device struct { + id string + name string + Type string + + IncomingCaps []string + OutgoingCaps []string + + state PairingState + CertFP string + + lastSeen time.Time + lastIP net.IP // cached from last successful connection; survives Disconnect + // lastPort is the tcpPort the peer last advertised over the authenticated + // (post-TLS) identity exchange. Used with lastIP as the dial target for + // paired devices so unauthenticated discovery packets can never redirect + // a paired auto-dial. Zero means unknown (fall back to 1716). + lastPort int + + // discoveryIP/discoveryPort remember where a device was last seen + // announcing itself (UDP/mDNS), even if we never opened a TCP + // connection to it. Used to dial on explicit user request + // (e.g. `kcd pair `) without auto-dialling strangers. + discoveryIP net.IP + discoveryPort int + + // pairDialRequested is the one-shot outbound dial trigger for an explicit + // `kcd pair ` request. It is consumed by the first discovery + // announcement after the request so the pair request can be delivered. + pairDialRequested atomic.Bool + + // pairIntentUntil is the Unix-nano deadline until which an explicit pair + // intent keeps a connection alive. Unlike pairDialRequested (consumed on + // first sighting), the intent survives dial/connect cycles until pairing + // starts, is rejected, succeeds, or the deadline (pairDialIntentTTL) + // expires — so a slow phone-side accept can't downgrade into an + // ephemeral-close flap. + pairIntentUntil atomic.Int64 + + // lastDiscoveryDial is when onDeviceFound last spawned a dial for this + // device. It throttles sighting-triggered redials so announcements + // (or a spoofed broadcast storm) can't cause a dial per packet. + lastDiscoveryDial time.Time + + // ephemeralDialed marks that this device already received its one + // ephemeral discovery dial for the current unpaired era. Ephemeral + // dials let a stranger complete the TCP identity exchange (so both + // sides list each other) without staying connected: the next sighting + // closes the socket again while the device is still unpaired. The + // marker is cleared when the device is explicitly unpaired/rejected, + // making it eligible again. Paired devices and pairing mode bypass it. + ephemeralDialed bool + + conn *transport.Conn + sendChan chan *protocol.Packet // buffered 32 + done chan struct{} + closeOnce sync.Once + + // lastConnect marks the last completed handshake; new handshakes + // inside reconnectCooldown are refused to starve duplicate bursts. + lastConnect time.Time + // lastSightedIP remembers the previous discovery sighting so only + // confirmed roams reset the reconnect backoff (see NoteSighting). + lastSightedIP net.IP + BatteryCharge int + IsCharging bool + + // batterySeen marks that at least one kdeconnect.battery packet was + // received. Until then the zero values above are not measurements — + // they must not be published (a fresh pair would otherwise report a + // stable, bogus 0% that no later packet corrects at steady charge). + batterySeen bool + // lastBatteryAt is when the last battery packet arrived, so clients + // can apply their own staleness rules (mirrors mediaAgeMs). + lastBatteryAt time.Time + + mu sync.RWMutex + + // reconnecting is an atomic flag preventing multiple concurrent + // auto-reconnect goroutines for this device. + reconnecting atomic.Bool + + // reconnectAttempt persists the auto-reconnect backoff counter across + // disconnect cycles. A connection that flaps (drops shortly after a + // successful dial) keeps the counter so the backoff escalates instead of + // resetting to the 2s floor; a stable connection resets it on drop. + reconnectAttempt int + + // connectStarted is when the most recent connection was established, + // used to detect flaps (connections that die too quickly to count as + // genuinely stable). + connectStarted time.Time + + // pluginDispatch routes incoming packets to registered plugins + pluginDispatch func(ctx context.Context, dev *Device, pkt *protocol.Packet) bool + onConnect func(dev *Device) + onDisconnect func(dev *Device) + + logger *zap.Logger + bus *events.Bus +} + +// NewDevice creates a new disconnected device instance. +// New devices start as Unpaired (not Unknown) so listings are unambiguous. +func NewDevice(id, name, dtype string, logger *zap.Logger) *Device { + return &Device{ + id: id, + name: name, + Type: dtype, + state: StateUnpaired, + sendChan: make(chan *protocol.Packet, 32), + done: make(chan struct{}), + logger: logger.With(zap.String("device_id", id)), + } +} + +// SetBus sets the event bus for the device. +func (d *Device) SetBus(bus *events.Bus) { + d.mu.Lock() + defer d.mu.Unlock() + d.bus = bus +} + +func (d *Device) ID() string { + return d.id +} +func (d *Device) Name() string { + d.mu.RLock() + defer d.mu.RUnlock() + return d.name +} +func (d *Device) SetName(n string) { + d.mu.Lock() + defer d.mu.Unlock() + d.name = n +} +func (d *Device) State() PairingState { + d.mu.RLock() + defer d.mu.RUnlock() + return d.state +} +func (d *Device) SetState(s PairingState) { + d.mu.Lock() + defer d.mu.Unlock() + d.state = s +} +func (d *Device) LastSeen() time.Time { + d.mu.RLock() + defer d.mu.RUnlock() + return d.lastSeen +} +func (d *Device) SetLastSeen(t time.Time) { + d.mu.Lock() + defer d.mu.Unlock() + d.lastSeen = t +} diff --git a/internal/device/device_discovery.go b/internal/device/device_discovery.go new file mode 100644 index 0000000..fdb7315 --- /dev/null +++ b/internal/device/device_discovery.go @@ -0,0 +1,193 @@ +package device + +import ( + "net" + "time" +) + +// SetDiscoveryAddr records where the device was last seen announcing itself. +func (d *Device) SetDiscoveryAddr(ip net.IP, port int) { + d.mu.Lock() + defer d.mu.Unlock() + d.discoveryIP = ip + d.discoveryPort = port +} + +// DiscoveryAddr returns the last-seen announcement address, or nil if unknown. +func (d *Device) DiscoveryAddr() (net.IP, int) { + d.mu.RLock() + defer d.mu.RUnlock() + return d.discoveryIP, d.discoveryPort +} + +// pairDialIntentTTL bounds how long an explicit `kcd pair ` intent pins +// a connection while pairing hasn't started yet. +const pairDialIntentTTL = 5 * time.Minute + +// RequestPairDial marks the device for a one-shot outbound dial on its next +// discovery announcement and arms the keep-alive intent until pairing starts, +// is rejected, succeeds, or the TTL expires. Used when the user explicitly +// runs `kcd pair ` for a device with no active connection. +func (d *Device) RequestPairDial(intentTTL ...time.Duration) { + ttl := pairDialIntentTTL + if len(intentTTL) > 0 && intentTTL[0] > 0 { + ttl = intentTTL[0] + } + d.pairDialRequested.Store(true) + d.pairIntentUntil.Store(time.Now().Add(ttl).UnixNano()) +} + +// ConsumePairDial reports and clears a pending explicit pair-dial request. +func (d *Device) ConsumePairDial() bool { + return d.pairDialRequested.CompareAndSwap(true, false) +} + +// PairDialPending reports whether an explicit pair-dial was requested, +// without clearing it. +func (d *Device) PairDialPending() bool { + return d.pairDialRequested.Load() +} + +// PairDialActive reports whether an explicit pair intent is still keeping +// the connection alive: requested and neither cleared nor expired. Unlike +// PairDialPending (the one-shot dial trigger, consumed on first sighting), +// this survives dial/connect cycles until pairing starts or ends. +func (d *Device) PairDialActive() bool { + return d.pairIntentUntil.Load() > time.Now().UnixNano() +} + +// ClearPairDial drops both the one-shot trigger and the keep-alive intent. +// Call it when pairing completes, is rejected/cancelled, or is unpaired. +func (d *Device) ClearPairDial() { + d.pairDialRequested.Store(false) + d.pairIntentUntil.Store(0) +} + +// LastPort returns the last authenticated tcpPort advertised by the peer, +// or 0 if unknown. +func (d *Device) LastPort() int { + d.mu.RLock() + defer d.mu.RUnlock() + return d.lastPort +} + +// SetLastPort records the peer's advertised listening port after a +// successful authenticated exchange. Callers must pass a validated port. +func (d *Device) SetLastPort(port int) { + d.mu.Lock() + defer d.mu.Unlock() + d.lastPort = port +} + +// SetLastIP records a dial target, used when restoring persisted state. +// A nil IP clears the target. +func (d *Device) SetLastIP(ip net.IP) { + d.mu.Lock() + defer d.mu.Unlock() + d.lastIP = ip +} + +// ShouldDiscoveryDial reports whether enough time has passed since the last +// discovery-triggered dial for this device, and marks this dial if so. It +// bounds redial storms to a stale LastIP (DHCP roam) or spoofed sightings. +func (d *Device) ShouldDiscoveryDial(minInterval time.Duration) bool { + d.mu.Lock() + defer d.mu.Unlock() + if time.Since(d.lastDiscoveryDial) < minInterval { + return false + } + d.lastDiscoveryDial = time.Now() + return true +} + +// MarkEphemeralDialed records that the one ephemeral discovery dial for the +// current unpaired era has been used. +func (d *Device) MarkEphemeralDialed() { + d.mu.Lock() + defer d.mu.Unlock() + d.ephemeralDialed = true +} + +// EphemeralDialed reports whether the ephemeral discovery dial was used. +func (d *Device) EphemeralDialed() bool { + d.mu.RLock() + defer d.mu.RUnlock() + return d.ephemeralDialed +} + +// ClearEphemeral makes the device eligible for a fresh ephemeral discovery +// dial (e.g. after an explicit unpair or rejection). +func (d *Device) ClearEphemeral() { + d.mu.Lock() + defer d.mu.Unlock() + d.ephemeralDialed = false +} + +// InCooldown reports whether a handshake completed too recently to start +// another one for this device. +func (d *Device) InCooldown() bool { + d.mu.RLock() + defer d.mu.RUnlock() + return !d.lastConnect.IsZero() && time.Since(d.lastConnect) < reconnectCooldown +} + +// NoteSighting records a discovery sighting while disconnected, reporting +// whether it confirms a genuine roam: same new address twice running. +// Single sightings prove nothing (IPv4/IPv6 alternation, AP flicker). +func (d *Device) NoteSighting(sighted net.IP) (roamed bool) { + d.mu.Lock() + defer d.mu.Unlock() + defer func() { d.lastSightedIP = sighted }() + if lastIP := d.lastIP; lastIP == nil || !lastIP.Equal(sighted) { + return sighted.Equal(d.lastSightedIP) + } + return false +} + +// TryReconnect attempts to mark the device as reconnecting. +// Returns true if this goroutine should proceed; false if another +// reconnect goroutine is already running. +func (d *Device) TryReconnect() bool { + return d.reconnecting.CompareAndSwap(false, true) +} + +// ReconnectDone marks the device as no longer reconnecting. +// Must be called (typically via defer) after a reconnect goroutine exits. +func (d *Device) ReconnectDone() { + d.reconnecting.Store(false) +} + +// ReconnectAttempt returns the persisted auto-reconnect backoff counter. +func (d *Device) ReconnectAttempt() int { + d.mu.RLock() + defer d.mu.RUnlock() + return d.reconnectAttempt +} + +// SetReconnectAttempt stores the auto-reconnect backoff counter so the next +// reconnect cycle (spawned when this connection drops) continues backing off +// instead of resetting to the initial floor. +func (d *Device) SetReconnectAttempt(n int) { + d.mu.Lock() + d.reconnectAttempt = n + d.mu.Unlock() +} + +// ResetReconnectAttempt clears the backoff counter after a stable connection +// (one that stayed up long enough to count as genuinely healthy) drops. +func (d *Device) ResetReconnectAttempt() { + d.mu.Lock() + d.reconnectAttempt = 0 + d.mu.Unlock() +} + +// ConnectionAge returns how long the current connection has been established, +// or 0 if the device has never connected in this process. +func (d *Device) ConnectionAge() time.Duration { + d.mu.RLock() + defer d.mu.RUnlock() + if d.connectStarted.IsZero() { + return 0 + } + return time.Since(d.connectStarted) +} diff --git a/internal/device/device_session.go b/internal/device/device_session.go new file mode 100644 index 0000000..a3af372 --- /dev/null +++ b/internal/device/device_session.go @@ -0,0 +1,176 @@ +package device + +import ( + "context" + "net" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// reconnectCooldown refuses new handshakes this long after a completed +// one, so post-roam bursts can't complete near-simultaneously and churn +// the peer's duplicate resolution. Reference stacks rate-limit the same +// way (desktop 500ms, Android MILLIS_DELAY_BETWEEN_CONNECTIONS_TO_SAME_DEVICE +// 1000ms); 1s matches upstream Android. +const reconnectCooldown = 1 * time.Second + +// Connect establishes a connection for the device and starts the reader and writer loops. +// A new authenticated connection immediately replaces any existing one +// (matching LanDeviceLink::reset / LanLink.reset). The old socket is +// closed; its readLoop will exit and disconnectConn will ignore it +// because d.conn no longer points at it. +func (d *Device) Connect(ctx context.Context, conn *transport.Conn, dispatch func(context.Context, *Device, *protocol.Packet) bool, onConnect func(*Device), onDisconnect func(*Device)) { + d.mu.Lock() + d.pluginDispatch = dispatch + d.onConnect = onConnect + d.onDisconnect = onDisconnect + + var oldConn *transport.Conn + var oldAddr, newAddr string + if d.conn != nil { + oldConn = d.conn + oldDone := d.done + oldAddr = oldConn.RemoteAddr().String() + newAddr = conn.RemoteAddr().String() + // Close the old done so its writerLoop exits; the new + // writerLoop will own the fresh channel. + d.closeOnce.Do(func() { + if oldDone != nil { + close(oldDone) + } + }) + } + + d.conn = conn + + // Renew the send channel on (re)connect in case it was closed during disconnect. + d.sendChan = make(chan *protocol.Packet, 32) + d.done = make(chan struct{}) + d.closeOnce = sync.Once{} + bus := d.bus + d.mu.Unlock() + + if oldConn != nil { + _ = oldConn.Close() + d.logger.Debug("replacing superseded connection", + zap.String("old_addr", oldAddr), + zap.String("new_addr", newAddr)) + } + + d.logger.Info("device connected", zap.String("remote_addr", conn.RemoteAddr().String())) + if bus != nil { + bus.Publish(events.TypeDeviceConnected, d.id, map[string]interface{}{ + "name": d.name, + "type": d.Type, + }) + } + + if d.onConnect != nil { + d.onConnect(d) + } + + // Cache the remote IP before the loops start so it is available + // after Disconnect() sets d.conn to nil (used by auto-reconnect). + if tcpAddr, ok := conn.RemoteAddr().(*net.TCPAddr); ok { + d.mu.Lock() + d.lastIP = tcpAddr.IP + d.mu.Unlock() + } + + // connectStarted distinguishes quick drops from stable connections, + // and doubles as the duplicate-cooldown clock (see InCooldown). + d.mu.Lock() + d.connectStarted = time.Now() + d.lastConnect = d.connectStarted + d.mu.Unlock() + + go d.readLoop(ctx, conn) + go d.writerLoop(ctx) +} + +// Disconnect terminates the session and stops the loops. +func (d *Device) Disconnect() { + d.mu.Lock() + d.lastSeen = time.Now() + + if d.conn == nil { + d.mu.Unlock() + return + } + + d.logger.Info("device disconnected") + _ = d.conn.Close() + d.conn = nil + d.lastConnect = time.Time{} + + // Capture these to call outside the lock to prevent deadlocks! + onDisc := d.onDisconnect + bus := d.bus + + d.closeOnce.Do(func() { + if d.done != nil { + close(d.done) + } + }) + d.mu.Unlock() + + // External calls must happen outside the mutex + if onDisc != nil { + onDisc(d) + } + if bus != nil { + bus.Publish(events.TypeDeviceDisconnected, d.id, nil) + } +} + +// IsConnected returns whether the device currently has an active connection. +func (d *Device) IsConnected() bool { + d.mu.RLock() + defer d.mu.RUnlock() + return d.conn != nil +} + +// disconnectConn handles a session's death. If the conn that died is no +// longer the current d.conn (it was superseded by a newer authenticated +// connection via Connect), the event is ignored silently — matching +// LanDeviceLink::reset's `if (m_socket == socket)` guard. Only the +// current preferred session's death triggers the full disconnect. +func (d *Device) disconnectConn(conn *transport.Conn) { + d.mu.Lock() + + if d.conn != conn { + d.mu.Unlock() + d.logger.Debug("ignoring disconnect from superseded connection") + return + } + + d.lastSeen = time.Now() + d.logger.Info("device disconnected") + _ = d.conn.Close() + d.conn = nil + d.lastConnect = time.Time{} + + // Capture callbacks to execute outside the lock + onDisc := d.onDisconnect + bus := d.bus + + d.closeOnce.Do(func() { + if d.done != nil { + close(d.done) + } + }) + d.mu.Unlock() + + // External calls must happen outside the mutex to prevent deadlocks + if onDisc != nil { + onDisc(d) + } + if bus != nil { + bus.Publish(events.TypeDeviceDisconnected, d.id, nil) + } +} diff --git a/internal/device/device_telemetry.go b/internal/device/device_telemetry.go new file mode 100644 index 0000000..8639523 --- /dev/null +++ b/internal/device/device_telemetry.go @@ -0,0 +1,108 @@ +package device + +import ( + "crypto/x509" + "net" + "time" + + "github.com/bethropolis/kcd/internal/events" +) + +// UpdateBattery updates the battery state of the device. +func (d *Device) UpdateBattery(charge int, charging bool) { + d.mu.Lock() + d.BatteryCharge = charge + d.IsCharging = charging + d.batterySeen = true + d.lastBatteryAt = time.Now() + bus := d.bus + id := d.id + d.mu.Unlock() + + if bus != nil { + bus.Publish(events.TypeBatteryUpdate, id, map[string]interface{}{ + "charge": charge, + "charging": charging, + }) + } +} + +// GetBattery returns the current battery state of the device. +func (d *Device) GetBattery() (int, bool) { + d.mu.RLock() + defer d.mu.RUnlock() + return d.BatteryCharge, d.IsCharging +} + +// HasBattery reports whether at least one battery packet was received. +// Until then the charge values are zero-value defaults, not measurements. +func (d *Device) HasBattery() bool { + d.mu.RLock() + defer d.mu.RUnlock() + return d.batterySeen +} + +// BatteryAge returns how long ago the last battery packet arrived, or a +// negative duration when no packet was ever received. +func (d *Device) BatteryAge() time.Duration { + d.mu.RLock() + defer d.mu.RUnlock() + if !d.batterySeen { + return -1 + } + return time.Since(d.lastBatteryAt) +} + +// HasCapability checks if the device has a particular capability (incoming or outgoing). +func (d *Device) HasCapability(cap string) bool { + d.mu.RLock() + defer d.mu.RUnlock() + + for _, c := range d.IncomingCaps { + if c == cap { + return true + } + } + for _, c := range d.OutgoingCaps { + if c == cap { + return true + } + } + return false +} + +// RemoteIP returns the IP address of the connected peer, if available. +func (d *Device) RemoteIP() net.IP { + d.mu.RLock() + defer d.mu.RUnlock() + + if d.conn == nil { + return nil + } + addr := d.conn.RemoteAddr() + if tcpAddr, ok := addr.(*net.TCPAddr); ok { + return tcpAddr.IP + } + return nil +} + +// LastIP returns the IP address from the most recent successful connection. +// Unlike RemoteIP, this persists after the connection drops — safe to use +// from an OnDisconnect callback for auto-reconnect dialling. +func (d *Device) LastIP() net.IP { + d.mu.RLock() + defer d.mu.RUnlock() + return d.lastIP +} + +// PeerCert returns the validated certificate presented by the remote device. +// Returns nil if not connected or no certificate was presented. +func (d *Device) PeerCert() *x509.Certificate { + d.mu.RLock() + defer d.mu.RUnlock() + + if d.conn == nil { + return nil + } + return d.conn.PeerCert() +} diff --git a/internal/discovery/broadcaster.go b/internal/discovery/broadcaster.go new file mode 100644 index 0000000..4685731 --- /dev/null +++ b/internal/discovery/broadcaster.go @@ -0,0 +1,101 @@ +package discovery + +import ( + "context" + "encoding/json" + "net" + "time" + + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// Broadcaster sends identity packets over UDP to advertise the local device. +type Broadcaster struct { + identityPacket *protocol.Packet + interval time.Duration + idleInterval time.Duration + logger *zap.Logger +} + +// NewBroadcaster creates a UDP discovery broadcaster. +func NewBroadcaster(identity *protocol.Packet, interval time.Duration, logger *zap.Logger) *Broadcaster { + return &Broadcaster{ + identityPacket: identity, + interval: interval, + logger: logger.With(zap.String("component", "broadcaster")), + } +} + +// Run periodically sends the identity packet to 255.255.255.255:1716. +// If shouldReduce is provided and returns true, the broadcast frequency +// is reduced to 60 seconds to save CPU and network resources while idle. +// (mDNS advertisement is no longer tied to this loop — see AdvertiseMDNS.) +func (b *Broadcaster) Run(ctx context.Context, shouldReduce func() bool) { + normalInterval := b.interval + reducedInterval := b.idleInterval + if reducedInterval <= 0 { + reducedInterval = 60 * time.Second + } + + conn, err := net.ListenUDP("udp4", nil) + if err != nil { + b.logger.Error("failed to listen for udp broadcast", zap.Error(err)) + return + } + defer conn.Close() + + data, err := json.Marshal(b.identityPacket) + if err != nil { + b.logger.Error("failed to marshal identity packet", zap.Error(err)) + return + } + data = append(data, '\n') + + timer := time.NewTimer(0) // Fire immediately on start + defer timer.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-timer.C: + // 1. Attempt global broadcast + globalAddr := &net.UDPAddr{IP: net.IPv4bcast, Port: 1716} + conn.WriteToUDP(data, globalAddr) + + // 2. Attempt per-interface directed broadcast for multi-homed reliability + ifaces, err := net.Interfaces() + if err == nil { + for _, iface := range ifaces { + if iface.Flags&net.FlagBroadcast == 0 || iface.Flags&net.FlagUp == 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + continue + } + for _, a := range addrs { + if ipnet, ok := a.(*net.IPNet); ok && ipnet.IP.To4() != nil { + ip4 := ipnet.IP.To4() + mask := ipnet.Mask[len(ipnet.Mask)-4:] + if len(mask) == 4 { + bcast := make(net.IP, 4) + for i := 0; i < 4; i++ { + bcast[i] = ip4[i] | ^mask[i] + } + conn.WriteToUDP(data, &net.UDPAddr{IP: bcast, Port: 1716}) + } + } + } + } + } + + nextInterval := normalInterval + if shouldReduce != nil && shouldReduce() { + nextInterval = reducedInterval + } + timer.Reset(nextInterval) + } + } +} diff --git a/internal/discovery/broadcaster_controller.go b/internal/discovery/broadcaster_controller.go new file mode 100644 index 0000000..8128e96 --- /dev/null +++ b/internal/discovery/broadcaster_controller.go @@ -0,0 +1,113 @@ +package discovery + +import ( + "context" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// Broadcast ownership: pairing mode (`kcd pair`) and the reconnect +// watcher share one loop but must not cancel each other, so starts are +// reference-counted per owner. The loop runs while any owner holds it. +const ( + OwnerPairing = "pairing" + OwnerReconnect = "reconnect" +) + +// BroadcasterController manages the broadcast lifecycle — start/stop on demand. +// Starts in stopped state. Broadcast is only active while Start() is in effect. +type BroadcasterController struct { + identityPacket *protocol.Packet + interval time.Duration + idleInterval time.Duration + shouldReduce func() bool + logger *zap.Logger + + mu sync.Mutex + running bool + cancel context.CancelFunc + owners map[string]struct{} +} + +// NewBroadcasterController creates a controller that starts in stopped state. +func NewBroadcasterController(identity *protocol.Packet, interval time.Duration, logger *zap.Logger, shouldReduce func() bool, idleInterval ...time.Duration) *BroadcasterController { + idle := 60 * time.Second + if len(idleInterval) > 0 && idleInterval[0] > 0 { + idle = idleInterval[0] + } + return &BroadcasterController{ + identityPacket: identity, + interval: interval, + idleInterval: idle, + shouldReduce: shouldReduce, + logger: logger.With(zap.String("component", "broadcaster")), + owners: make(map[string]struct{}), + } +} + +// Start launches the UDP broadcaster loop for the pairing owner. +// No-op if already running (ownership is still recorded). +func (bc *BroadcasterController) Start(parentCtx context.Context) { + bc.StartOwned(parentCtx, OwnerPairing) +} + +// Stop withdraws the pairing owner. The loop stops only when no owners +// remain, so a reconnect-driven broadcast survives `kcd pair` exiting. +func (bc *BroadcasterController) Stop() { + bc.StopOwned(OwnerPairing) +} + +// StartOwned launches the loop (if needed) and records owner as needing it. +func (bc *BroadcasterController) StartOwned(parentCtx context.Context, owner string) { + bc.mu.Lock() + defer bc.mu.Unlock() + + bc.owners[owner] = struct{}{} + if bc.running { + return + } + + ctx, cancel := context.WithCancel(parentCtx) + bc.cancel = cancel + bc.running = true + + b := &Broadcaster{ + identityPacket: bc.identityPacket, + interval: bc.interval, + idleInterval: bc.idleInterval, + logger: bc.logger, + } + go func() { + b.Run(ctx, bc.shouldReduce) + bc.mu.Lock() + bc.running = false + bc.mu.Unlock() + }() +} + +// StopOwned withdraws owner's need. No-op if the owner holds nothing. +func (bc *BroadcasterController) StopOwned(owner string) { + bc.mu.Lock() + defer bc.mu.Unlock() + + if _, ok := bc.owners[owner]; !ok { + return + } + delete(bc.owners, owner) + if len(bc.owners) > 0 || !bc.running || bc.cancel == nil { + return + } + bc.cancel() + bc.cancel = nil + bc.running = false +} + +// IsRunning reports whether the broadcast loop is currently active. +func (bc *BroadcasterController) IsRunning() bool { + bc.mu.Lock() + defer bc.mu.Unlock() + return bc.running +} diff --git a/internal/discovery/discovery.go b/internal/discovery/discovery.go deleted file mode 100644 index 67e7a03..0000000 --- a/internal/discovery/discovery.go +++ /dev/null @@ -1,380 +0,0 @@ -package discovery - -import ( - "context" - "encoding/json" - "net" - "strconv" - "strings" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/protocol" - "github.com/libp2p/zeroconf/v2" - "go.uber.org/zap" -) - -// Broadcast ownership: pairing mode (`kcd pair`) and the reconnect -// watcher share one loop but must not cancel each other, so starts are -// reference-counted per owner. The loop runs while any owner holds it. -const ( - OwnerPairing = "pairing" - OwnerReconnect = "reconnect" -) - -// BroadcasterController manages the broadcast lifecycle — start/stop on demand. -// Starts in stopped state. Broadcast is only active while Start() is in effect. -type BroadcasterController struct { - identityPacket *protocol.Packet - interval time.Duration - idleInterval time.Duration - shouldReduce func() bool - logger *zap.Logger - - mu sync.Mutex - running bool - cancel context.CancelFunc - owners map[string]struct{} -} - -// NewBroadcasterController creates a controller that starts in stopped state. -func NewBroadcasterController(identity *protocol.Packet, interval time.Duration, logger *zap.Logger, shouldReduce func() bool, idleInterval ...time.Duration) *BroadcasterController { - idle := 60 * time.Second - if len(idleInterval) > 0 && idleInterval[0] > 0 { - idle = idleInterval[0] - } - return &BroadcasterController{ - identityPacket: identity, - interval: interval, - idleInterval: idle, - shouldReduce: shouldReduce, - logger: logger.With(zap.String("component", "broadcaster")), - owners: make(map[string]struct{}), - } -} - -// Start launches the UDP broadcaster loop for the pairing owner. -// No-op if already running (ownership is still recorded). -func (bc *BroadcasterController) Start(parentCtx context.Context) { - bc.StartOwned(parentCtx, OwnerPairing) -} - -// Stop withdraws the pairing owner. The loop stops only when no owners -// remain, so a reconnect-driven broadcast survives `kcd pair` exiting. -func (bc *BroadcasterController) Stop() { - bc.StopOwned(OwnerPairing) -} - -// StartOwned launches the loop (if needed) and records owner as needing it. -func (bc *BroadcasterController) StartOwned(parentCtx context.Context, owner string) { - bc.mu.Lock() - defer bc.mu.Unlock() - - bc.owners[owner] = struct{}{} - if bc.running { - return - } - - ctx, cancel := context.WithCancel(parentCtx) - bc.cancel = cancel - bc.running = true - - b := &Broadcaster{ - identityPacket: bc.identityPacket, - interval: bc.interval, - idleInterval: bc.idleInterval, - logger: bc.logger, - } - go func() { - b.Run(ctx, bc.shouldReduce) - bc.mu.Lock() - bc.running = false - bc.mu.Unlock() - }() -} - -// StopOwned withdraws owner's need. No-op if the owner holds nothing. -func (bc *BroadcasterController) StopOwned(owner string) { - bc.mu.Lock() - defer bc.mu.Unlock() - - if _, ok := bc.owners[owner]; !ok { - return - } - delete(bc.owners, owner) - if len(bc.owners) > 0 || !bc.running || bc.cancel == nil { - return - } - bc.cancel() - bc.cancel = nil - bc.running = false -} - -// IsRunning reports whether the broadcast loop is currently active. -func (bc *BroadcasterController) IsRunning() bool { - bc.mu.Lock() - defer bc.mu.Unlock() - return bc.running -} - -// AdvertiseMDNS registers the local identity as _kdeconnect._udp until -// ctx ends. Unlike UDP broadcast this is responder-only (it wakes on -// incoming queries), so it stays up for the daemon lifetime at negligible -// idle cost and gives phones a standing discovery path even while UDP -// broadcast is stopped. -func AdvertiseMDNS(ctx context.Context, identityPacket *protocol.Packet, logger *zap.Logger) { - var idBody protocol.IdentityBody - if err := json.Unmarshal(identityPacket.Body, &idBody); err != nil { - logger.Warn("failed to parse identity for mDNS", zap.Error(err)) - return - } - server, err := zeroconf.Register( - idBody.DeviceName, - "_kdeconnect._udp", - "local.", - idBody.TCPPort, - []string{ - "id=" + idBody.DeviceID, - "name=" + idBody.DeviceName, - "type=" + idBody.DeviceType, - "protocol=8", - }, - nil, - ) - if err != nil { - logger.Warn("failed to register mDNS service", zap.Error(err)) - return - } - go func() { - <-ctx.Done() - server.Shutdown() - logger.Info("mDNS service shut down") - }() -} - -// Broadcaster sends identity packets over UDP to advertise the local device. -type Broadcaster struct { - identityPacket *protocol.Packet - interval time.Duration - idleInterval time.Duration - logger *zap.Logger -} - -// NewBroadcaster creates a UDP discovery broadcaster. -func NewBroadcaster(identity *protocol.Packet, interval time.Duration, logger *zap.Logger) *Broadcaster { - return &Broadcaster{ - identityPacket: identity, - interval: interval, - logger: logger.With(zap.String("component", "broadcaster")), - } -} - -// Run periodically sends the identity packet to 255.255.255.255:1716. -// If shouldReduce is provided and returns true, the broadcast frequency -// is reduced to 60 seconds to save CPU and network resources while idle. -// (mDNS advertisement is no longer tied to this loop — see AdvertiseMDNS.) -func (b *Broadcaster) Run(ctx context.Context, shouldReduce func() bool) { - normalInterval := b.interval - reducedInterval := b.idleInterval - if reducedInterval <= 0 { - reducedInterval = 60 * time.Second - } - - conn, err := net.ListenUDP("udp4", nil) - if err != nil { - b.logger.Error("failed to listen for udp broadcast", zap.Error(err)) - return - } - defer conn.Close() - - data, err := json.Marshal(b.identityPacket) - if err != nil { - b.logger.Error("failed to marshal identity packet", zap.Error(err)) - return - } - data = append(data, '\n') - - timer := time.NewTimer(0) // Fire immediately on start - defer timer.Stop() - - for { - select { - case <-ctx.Done(): - return - case <-timer.C: - // 1. Attempt global broadcast - globalAddr := &net.UDPAddr{IP: net.IPv4bcast, Port: 1716} - conn.WriteToUDP(data, globalAddr) - - // 2. Attempt per-interface directed broadcast for multi-homed reliability - ifaces, err := net.Interfaces() - if err == nil { - for _, iface := range ifaces { - if iface.Flags&net.FlagBroadcast == 0 || iface.Flags&net.FlagUp == 0 { - continue - } - addrs, err := iface.Addrs() - if err != nil { - continue - } - for _, a := range addrs { - if ipnet, ok := a.(*net.IPNet); ok && ipnet.IP.To4() != nil { - ip4 := ipnet.IP.To4() - mask := ipnet.Mask[len(ipnet.Mask)-4:] - if len(mask) == 4 { - bcast := make(net.IP, 4) - for i := 0; i < 4; i++ { - bcast[i] = ip4[i] | ^mask[i] - } - conn.WriteToUDP(data, &net.UDPAddr{IP: bcast, Port: 1716}) - } - } - } - } - } - - nextInterval := normalInterval - if shouldReduce != nil && shouldReduce() { - nextInterval = reducedInterval - } - timer.Reset(nextInterval) - } - } -} - -// Listener listens for UDP identity packets from other devices. -type Listener struct { - port int - localDeviceID string - onDeviceFound func(ip net.IP, tcpPort int, identity *protocol.Packet) - logger *zap.Logger -} - -// NewListener creates a UDP discovery listener. -func NewListener(port int, localDeviceID string, callback func(ip net.IP, tcpPort int, identity *protocol.Packet), logger *zap.Logger) *Listener { - return &Listener{ - port: port, - localDeviceID: localDeviceID, - onDeviceFound: callback, - logger: logger.With(zap.String("component", "udp-listener")), - } -} - -// Run starts the UDP listener loop to parse incoming discovery broadcasts. -func (l *Listener) Run(ctx context.Context) { - // mDNS Discovery - go l.runMdnsDiscovery(ctx) - - addr := &net.UDPAddr{Port: l.port} - conn, err := net.ListenUDP("udp", addr) - if err != nil { - l.logger.Error("failed to listen on udp", zap.Int("port", l.port), zap.Error(err)) - return - } - defer conn.Close() - - // 8KB buffer shouldn't be exceeded by an identity packet - buf := make([]byte, 8192) - - go func() { - <-ctx.Done() - conn.Close() - }() - - for { - n, remoteAddr, err := conn.ReadFromUDP(buf) - if err != nil { - if ctx.Err() != nil { - return // clean exit on context cancel - } - l.logger.Debug("read udp error", zap.Error(err)) - continue - } - - if n >= len(buf) || n == 0 { - continue // ignore giant/empty packets - } - - pkt := protocol.AcquirePacket() - if err := json.Unmarshal(buf[:n], pkt); err != nil { - protocol.ReleasePacket(pkt) - continue - } - - if pkt.Type != protocol.TypeIdentity { - protocol.ReleasePacket(pkt) - continue - } - - var identity protocol.IdentityBody - if err := json.Unmarshal(pkt.Body, &identity); err != nil { - protocol.ReleasePacket(pkt) - continue - } - - // Don't connect to ourselves - if identity.DeviceID == l.localDeviceID { - protocol.ReleasePacket(pkt) - continue - } - - if l.onDeviceFound != nil { - l.onDeviceFound(remoteAddr.IP, identity.TCPPort, pkt) - } - protocol.ReleasePacket(pkt) - } -} - -func (l *Listener) runMdnsDiscovery(ctx context.Context) { - entries := make(chan *zeroconf.ServiceEntry) - go func(results <-chan *zeroconf.ServiceEntry) { - for entry := range results { - var deviceId, deviceName, deviceType string - var protocolVersion int - for _, txt := range entry.Text { - key, val, ok := strings.Cut(txt, "=") - if !ok { - continue - } - switch key { - case "id": - deviceId = val - case "name": - deviceName = val - case "type": - deviceType = val - case "protocol": - protocolVersion, _ = strconv.Atoi(val) - } - } - - if deviceId == "" || deviceId == l.localDeviceID { - continue - } - if len(entry.AddrIPv4) == 0 { - continue - } - - body := protocol.IdentityBody{ - DeviceID: deviceId, - DeviceName: deviceName, - DeviceType: deviceType, - ProtocolVersion: protocolVersion, - TCPPort: entry.Port, - } - pkt, err := protocol.NewPacket(protocol.TypeIdentity, body) - if err != nil { - continue - } - - if l.onDeviceFound != nil { - l.onDeviceFound(entry.AddrIPv4[0], entry.Port, pkt) - } - protocol.ReleasePacket(pkt) - } - }(entries) - - if err := zeroconf.Browse(ctx, "_kdeconnect._udp", "local.", entries); err != nil { - l.logger.Warn("failed to browse mDNS", zap.Error(err)) - } -} diff --git a/internal/discovery/listener.go b/internal/discovery/listener.go new file mode 100644 index 0000000..feea722 --- /dev/null +++ b/internal/discovery/listener.go @@ -0,0 +1,93 @@ +package discovery + +import ( + "context" + "encoding/json" + "net" + + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// Listener listens for UDP identity packets from other devices. +type Listener struct { + port int + localDeviceID string + onDeviceFound func(ip net.IP, tcpPort int, identity *protocol.Packet) + logger *zap.Logger +} + +// NewListener creates a UDP discovery listener. +func NewListener(port int, localDeviceID string, callback func(ip net.IP, tcpPort int, identity *protocol.Packet), logger *zap.Logger) *Listener { + return &Listener{ + port: port, + localDeviceID: localDeviceID, + onDeviceFound: callback, + logger: logger.With(zap.String("component", "udp-listener")), + } +} + +// Run starts the UDP listener loop to parse incoming discovery broadcasts. +func (l *Listener) Run(ctx context.Context) { + // mDNS Discovery + go l.runMdnsDiscovery(ctx) + + addr := &net.UDPAddr{Port: l.port} + conn, err := net.ListenUDP("udp", addr) + if err != nil { + l.logger.Error("failed to listen on udp", zap.Int("port", l.port), zap.Error(err)) + return + } + defer conn.Close() + + // 8KB buffer shouldn't be exceeded by an identity packet + buf := make([]byte, 8192) + + go func() { + <-ctx.Done() + conn.Close() + }() + + for { + n, remoteAddr, err := conn.ReadFromUDP(buf) + if err != nil { + if ctx.Err() != nil { + return // clean exit on context cancel + } + l.logger.Debug("read udp error", zap.Error(err)) + continue + } + + if n >= len(buf) || n == 0 { + continue // ignore giant/empty packets + } + + pkt := protocol.AcquirePacket() + if err := json.Unmarshal(buf[:n], pkt); err != nil { + protocol.ReleasePacket(pkt) + continue + } + + if pkt.Type != protocol.TypeIdentity { + protocol.ReleasePacket(pkt) + continue + } + + var identity protocol.IdentityBody + if err := json.Unmarshal(pkt.Body, &identity); err != nil { + protocol.ReleasePacket(pkt) + continue + } + + // Don't connect to ourselves + if identity.DeviceID == l.localDeviceID { + protocol.ReleasePacket(pkt) + continue + } + + if l.onDeviceFound != nil { + l.onDeviceFound(remoteAddr.IP, identity.TCPPort, pkt) + } + protocol.ReleasePacket(pkt) + } +} diff --git a/internal/discovery/listener_mdns.go b/internal/discovery/listener_mdns.go new file mode 100644 index 0000000..c5a699f --- /dev/null +++ b/internal/discovery/listener_mdns.go @@ -0,0 +1,105 @@ +package discovery + +import ( + "context" + "encoding/json" + "strconv" + "strings" + + "github.com/bethropolis/kcd/internal/protocol" + "github.com/libp2p/zeroconf/v2" + "go.uber.org/zap" +) + +// AdvertiseMDNS registers the local identity as _kdeconnect._udp until +// ctx ends. Unlike UDP broadcast this is responder-only (it wakes on +// incoming queries), so it stays up for the daemon lifetime at negligible +// idle cost and gives phones a standing discovery path even while UDP +// broadcast is stopped. +func AdvertiseMDNS(ctx context.Context, identityPacket *protocol.Packet, logger *zap.Logger) { + var idBody protocol.IdentityBody + if err := json.Unmarshal(identityPacket.Body, &idBody); err != nil { + logger.Warn("failed to parse identity for mDNS", zap.Error(err)) + return + } + server, err := zeroconf.Register( + idBody.DeviceName, + "_kdeconnect._udp", + "local.", + idBody.TCPPort, + []string{ + "id=" + idBody.DeviceID, + "name=" + idBody.DeviceName, + "type=" + idBody.DeviceType, + "protocol=8", + }, + nil, + ) + if err != nil { + logger.Warn("failed to register mDNS service", zap.Error(err)) + return + } + go func() { + <-ctx.Done() + server.Shutdown() + logger.Info("mDNS service shut down") + }() +} + +// runMdnsDiscovery browses for peer _kdeconnect._udp services and feeds +// sightings to onDeviceFound as synthetic identity packets. It is a method +// on Listener (kept apart from the UDP Run loop) so both transports share +// the same callback and self-filtering. +func (l *Listener) runMdnsDiscovery(ctx context.Context) { + entries := make(chan *zeroconf.ServiceEntry) + go func(results <-chan *zeroconf.ServiceEntry) { + for entry := range results { + var deviceId, deviceName, deviceType string + var protocolVersion int + for _, txt := range entry.Text { + key, val, ok := strings.Cut(txt, "=") + if !ok { + continue + } + switch key { + case "id": + deviceId = val + case "name": + deviceName = val + case "type": + deviceType = val + case "protocol": + protocolVersion, _ = strconv.Atoi(val) + } + } + + if deviceId == "" || deviceId == l.localDeviceID { + continue + } + if len(entry.AddrIPv4) == 0 { + continue + } + + body := protocol.IdentityBody{ + DeviceID: deviceId, + DeviceName: deviceName, + DeviceType: deviceType, + ProtocolVersion: protocolVersion, + TCPPort: entry.Port, + } + pkt, err := protocol.NewPacket(protocol.TypeIdentity, body) + if err != nil { + continue + } + + if l.onDeviceFound != nil { + l.onDeviceFound(entry.AddrIPv4[0], entry.Port, pkt) + } + protocol.ReleasePacket(pkt) + } + }(entries) + + if err := zeroconf.Browse(ctx, "_kdeconnect._udp", "local.", entries); err != nil { + l.logger.Warn("failed to browse mDNS", zap.Error(err)) + } +} diff --git a/internal/ipc/server.go b/internal/ipc/server.go index d28b9ca..9e511f7 100644 --- a/internal/ipc/server.go +++ b/internal/ipc/server.go @@ -9,12 +9,8 @@ import ( "os" "path/filepath" "strconv" - "time" "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/plugins/connectivity" - "github.com/bethropolis/kcd/internal/plugins/mpris" "go.uber.org/zap" ) @@ -154,145 +150,3 @@ func (s *Server) writeResponse(conn net.Conn, res Response) { data = append(data, '\n') _, _ = conn.Write(data) } - -func (s *Server) handleWatch(conn net.Conn, payload []byte) { - defer conn.Close() - - var p WatchPayload - if len(payload) > 0 { - _ = json.Unmarshal(payload, &p) - } - - bus := s.handler.bus - if bus == nil { - s.writeResponse(conn, Response{OK: false, Error: "event bus not enabled"}) - return - } - - // Send OK response to indicate stream is starting - s.writeResponse(conn, Response{OK: true}) - - // Full-state snapshot first: every known device (online AND offline) - // with cached battery/media/signal, so clients boot with complete - // state from this single connection — no auxiliary bootstrap calls, - // no hydration races. Sent regardless of event filters. - snapEv := map[string]interface{}{ - "type": events.TypeStateSnapshot, - "timestamp": time.Now().UTC(), - "payload": BuildSnapshot(s.handler.devices, s.handler.plugins), - } - data, _ := json.Marshal(snapEv) - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return - } - - // Initial State Dump - // For each connected device, we emit device.connected and battery.update. - devs := s.handler.devices.Connected() - for _, dev := range devs { - devData := map[string]interface{}{ - "id": dev.ID(), - "name": dev.Name(), - "type": dev.Type, - } - - // Send initial connected event - initEv := map[string]interface{}{ - "type": "device.connected", - "deviceId": dev.ID(), - "timestamp": time.Now().UTC(), - "payload": devData, - } - data, _ = json.Marshal(initEv) - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return - } - - // Send initial battery event, but only when the daemon actually - // has a reading. Emitting zero values for a fresh pair would - // publish a bogus stable 0% (see Device.HasBattery) — same - // skip-if-absent rule as connectivity below. - if dev.HasBattery() { - charge, charging := dev.GetBattery() - batEv := map[string]interface{}{ - "type": "battery.update", - "deviceId": dev.ID(), - "timestamp": time.Now().UTC(), - "payload": map[string]interface{}{ - "charge": charge, - "charging": charging, - "batteryAgeMs": dev.BatteryAge().Milliseconds(), - }, - } - data, _ = json.Marshal(batEv) - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return - } - } - - // Send initial connectivity state if the device already reported. - // Unlike battery there is no meaningful zero value, so devices - // without a report are skipped instead of emitting empty data. - if pl, ok := s.handler.plugins.GetByName("Connectivity"); ok { - if report, ok := pl.(*connectivity.ConnectivityPlugin).Report(dev.ID()); ok { - connEv := map[string]interface{}{ - "type": "connectivity.update", - "deviceId": dev.ID(), - "timestamp": time.Now().UTC(), - "payload": report, - } - data, _ = json.Marshal(connEv) - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return - } - } - } - - // Send initial mpris state if recently updated. - // If the last update was more than 10 seconds ago the state is - // likely stale — the phone probably stopped playing — so we skip it - // to avoid showing a ghost "now playing" in Waybar/CLI after reconnect. - if pl, ok := s.handler.plugins.GetByName("MPRIS"); ok { - mp := pl.(*mpris.MPRISPlugin) - if mp.RemoteStateAge(dev.ID()) < 10*time.Second { - if state := mp.RemoteState(dev.ID()); state != nil { - mprisEv := map[string]interface{}{ - "type": "mpris.update", - "deviceId": dev.ID(), - "timestamp": time.Now().UTC(), - "payload": state, - } - data, _ = json.Marshal(mprisEv) - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return - } - } - } - } - } - - // Subscribe - var filters []events.EventType - for _, e := range p.Events { - filters = append(filters, events.EventType(e)) - } - - sub := bus.Subscribe(events.WatchSubscriberCap, filters...) - defer sub.Close() - - for ev := range sub.C { - data, err := json.Marshal(ev) - if err != nil { - continue - } - data = append(data, '\n') - if _, err := conn.Write(data); err != nil { - return // connection likely closed - } - } -} diff --git a/internal/ipc/server_watch.go b/internal/ipc/server_watch.go new file mode 100644 index 0000000..6f6ec89 --- /dev/null +++ b/internal/ipc/server_watch.go @@ -0,0 +1,153 @@ +package ipc + +import ( + "encoding/json" + "net" + "time" + + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/plugins/connectivity" + "github.com/bethropolis/kcd/internal/plugins/mpris" +) + +func (s *Server) handleWatch(conn net.Conn, payload []byte) { + defer conn.Close() + + var p WatchPayload + if len(payload) > 0 { + _ = json.Unmarshal(payload, &p) + } + + bus := s.handler.bus + if bus == nil { + s.writeResponse(conn, Response{OK: false, Error: "event bus not enabled"}) + return + } + + // Send OK response to indicate stream is starting + s.writeResponse(conn, Response{OK: true}) + + // Full-state snapshot first: every known device (online AND offline) + // with cached battery/media/signal, so clients boot with complete + // state from this single connection — no auxiliary bootstrap calls, + // no hydration races. Sent regardless of event filters. + snapEv := map[string]interface{}{ + "type": events.TypeStateSnapshot, + "timestamp": time.Now().UTC(), + "payload": BuildSnapshot(s.handler.devices, s.handler.plugins), + } + data, _ := json.Marshal(snapEv) + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return + } + + // Initial State Dump + // For each connected device, we emit device.connected and battery.update. + devs := s.handler.devices.Connected() + for _, dev := range devs { + devData := map[string]interface{}{ + "id": dev.ID(), + "name": dev.Name(), + "type": dev.Type, + } + + // Send initial connected event + initEv := map[string]interface{}{ + "type": "device.connected", + "deviceId": dev.ID(), + "timestamp": time.Now().UTC(), + "payload": devData, + } + data, _ = json.Marshal(initEv) + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return + } + + // Send initial battery event, but only when the daemon actually + // has a reading. Emitting zero values for a fresh pair would + // publish a bogus stable 0% (see Device.HasBattery) — same + // skip-if-absent rule as connectivity below. + if dev.HasBattery() { + charge, charging := dev.GetBattery() + batEv := map[string]interface{}{ + "type": "battery.update", + "deviceId": dev.ID(), + "timestamp": time.Now().UTC(), + "payload": map[string]interface{}{ + "charge": charge, + "charging": charging, + "batteryAgeMs": dev.BatteryAge().Milliseconds(), + }, + } + data, _ = json.Marshal(batEv) + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return + } + } + + // Send initial connectivity state if the device already reported. + // Unlike battery there is no meaningful zero value, so devices + // without a report are skipped instead of emitting empty data. + if pl, ok := s.handler.plugins.GetByName("Connectivity"); ok { + if report, ok := pl.(*connectivity.ConnectivityPlugin).Report(dev.ID()); ok { + connEv := map[string]interface{}{ + "type": "connectivity.update", + "deviceId": dev.ID(), + "timestamp": time.Now().UTC(), + "payload": report, + } + data, _ = json.Marshal(connEv) + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return + } + } + } + + // Send initial mpris state if recently updated. + // If the last update was more than 10 seconds ago the state is + // likely stale — the phone probably stopped playing — so we skip it + // to avoid showing a ghost "now playing" in Waybar/CLI after reconnect. + if pl, ok := s.handler.plugins.GetByName("MPRIS"); ok { + mp := pl.(*mpris.MPRISPlugin) + if mp.RemoteStateAge(dev.ID()) < 10*time.Second { + if state := mp.RemoteState(dev.ID()); state != nil { + mprisEv := map[string]interface{}{ + "type": "mpris.update", + "deviceId": dev.ID(), + "timestamp": time.Now().UTC(), + "payload": state, + } + data, _ = json.Marshal(mprisEv) + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return + } + } + } + } + } + + // Subscribe + var filters []events.EventType + for _, e := range p.Events { + filters = append(filters, events.EventType(e)) + } + + sub := bus.Subscribe(events.WatchSubscriberCap, filters...) + defer sub.Close() + + for ev := range sub.C { + data, err := json.Marshal(ev) + if err != nil { + continue + } + data = append(data, '\n') + if _, err := conn.Write(data); err != nil { + return // connection likely closed + } + } +} diff --git a/internal/plugins/battery/battery.go b/internal/plugins/battery/handle.go similarity index 50% rename from internal/plugins/battery/battery.go rename to internal/plugins/battery/handle.go index dcebd9f..39c73ab 100644 --- a/internal/plugins/battery/battery.go +++ b/internal/plugins/battery/handle.go @@ -4,13 +4,7 @@ import ( "context" "encoding/json" "fmt" - "os" - "path/filepath" - "strconv" - "strings" - "time" - "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/plugin" @@ -18,54 +12,6 @@ import ( "go.uber.org/zap" ) -// ThresholdEvent values from the KDE Connect protocol. -const ( - thresholdNone = 0 - thresholdLow = 1 // battery is low (typically <= 15%) - thresholdFull = 2 // battery reached full charge -) - -// BatteryPlugin handles incoming battery state updates. -type BatteryPlugin struct { - notifications config.NotificationConfig - cfg config.BatteryConfig - bus *events.Bus - logger *zap.Logger -} - -// NewBatteryPlugin creates a BatteryPlugin. -func NewBatteryPlugin(cfg config.BatteryConfig, bus *events.Bus, logger *zap.Logger, notifications ...config.NotificationConfig) *BatteryPlugin { - var notificationCfg config.NotificationConfig - if len(notifications) > 0 { - notificationCfg = notifications[0] - } - return &BatteryPlugin{ - notifications: notificationCfg, - cfg: cfg, - bus: bus, - logger: logger.With(zap.String("plugin", "battery")), - } -} - -// BatteryBody represents the body of a kdeconnect.battery packet. -// A body carrying only Request asks the peer to report its state; it is -// not a state update and must never touch the stored charge. -type BatteryBody struct { - CurrentCharge int `json:"currentCharge"` - IsCharging bool `json:"isCharging"` - ThresholdEvent int `json:"thresholdEvent"` - Request bool `json:"request,omitempty"` -} - -func (p *BatteryPlugin) Name() string { return "Battery" } -func (p *BatteryPlugin) Timeout() time.Duration { return 5 * time.Second } -func (p *BatteryPlugin) IncomingTypes() []string { - return []string{"kdeconnect.battery", "kdeconnect.battery.request"} -} -func (p *BatteryPlugin) OutgoingTypes() []string { - return []string{"kdeconnect.battery", "kdeconnect.battery.request"} -} - // Handle processes incoming battery packets. func (p *BatteryPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { switch pkt.Type { @@ -177,58 +123,3 @@ func (p *BatteryPlugin) OnConnect(dev device.Sender) { } func (p *BatteryPlugin) OnDisconnect(_ device.Sender) {} - -// powerSupplyRoots lists paths to check for battery sysfs entries. -var powerSupplyRoots = []string{ - "/sys/class/power_supply", - "/sys/devices/platform/subsystem/power_supply", -} - -// readLocalBattery reads the local battery state from sysfs. -// Returns charge (0-100), charging status, and any error. -// If no battery is found, returns an error — callers should log and skip. -func readLocalBattery() (int, bool, error) { - for _, root := range powerSupplyRoots { - entries, err := os.ReadDir(root) - if err != nil { - continue - } - for _, e := range entries { - name := e.Name() - if !strings.HasPrefix(name, "BAT") { - continue - } - base := filepath.Join(root, name) - - // Read capacity (0-100) - capRaw, err := os.ReadFile(filepath.Join(base, "capacity")) - if err != nil { - continue - } - capacity, err := strconv.Atoi(strings.TrimSpace(string(capRaw))) - if err != nil { - continue - } - - // Read status (Charging/Discharging/Full/Unknown) - statusRaw, _ := os.ReadFile(filepath.Join(base, "status")) - status := strings.TrimSpace(string(statusRaw)) - charging := status == "Charging" - - return capacity, charging, nil - } - } - - return 0, false, errNoBattery -} - -// errNoBattery is returned when no battery sysfs entry is found. -var errNoBattery = &noBatteryError{} - -type noBatteryError struct{} - -func (e *noBatteryError) Error() string { return "no battery found" } -func (e *noBatteryError) Is(target error) bool { - _, ok := target.(*noBatteryError) - return ok -} diff --git a/internal/plugins/battery/local.go b/internal/plugins/battery/local.go new file mode 100644 index 0000000..41b2747 --- /dev/null +++ b/internal/plugins/battery/local.go @@ -0,0 +1,63 @@ +package battery + +import ( + "os" + "path/filepath" + "strconv" + "strings" +) + +// powerSupplyRoots lists paths to check for battery sysfs entries. +var powerSupplyRoots = []string{ + "/sys/class/power_supply", + "/sys/devices/platform/subsystem/power_supply", +} + +// readLocalBattery reads the local battery state from sysfs. +// Returns charge (0-100), charging status, and any error. +// If no battery is found, returns an error — callers should log and skip. +func readLocalBattery() (int, bool, error) { + for _, root := range powerSupplyRoots { + entries, err := os.ReadDir(root) + if err != nil { + continue + } + for _, e := range entries { + name := e.Name() + if !strings.HasPrefix(name, "BAT") { + continue + } + base := filepath.Join(root, name) + + // Read capacity (0-100) + capRaw, err := os.ReadFile(filepath.Join(base, "capacity")) + if err != nil { + continue + } + capacity, err := strconv.Atoi(strings.TrimSpace(string(capRaw))) + if err != nil { + continue + } + + // Read status (Charging/Discharging/Full/Unknown) + statusRaw, _ := os.ReadFile(filepath.Join(base, "status")) + status := strings.TrimSpace(string(statusRaw)) + charging := status == "Charging" + + return capacity, charging, nil + } + } + + return 0, false, errNoBattery +} + +// errNoBattery is returned when no battery sysfs entry is found. +var errNoBattery = &noBatteryError{} + +type noBatteryError struct{} + +func (e *noBatteryError) Error() string { return "no battery found" } +func (e *noBatteryError) Is(target error) bool { + _, ok := target.(*noBatteryError) + return ok +} diff --git a/internal/plugins/battery/types.go b/internal/plugins/battery/types.go new file mode 100644 index 0000000..4c62d33 --- /dev/null +++ b/internal/plugins/battery/types.go @@ -0,0 +1,57 @@ +package battery + +import ( + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/events" + "go.uber.org/zap" +) + +// ThresholdEvent values from the KDE Connect protocol. +const ( + thresholdNone = 0 + thresholdLow = 1 // battery is low (typically <= 15%) + thresholdFull = 2 // battery reached full charge +) + +// BatteryPlugin handles incoming battery state updates. +type BatteryPlugin struct { + notifications config.NotificationConfig + cfg config.BatteryConfig + bus *events.Bus + logger *zap.Logger +} + +// NewBatteryPlugin creates a BatteryPlugin. +func NewBatteryPlugin(cfg config.BatteryConfig, bus *events.Bus, logger *zap.Logger, notifications ...config.NotificationConfig) *BatteryPlugin { + var notificationCfg config.NotificationConfig + if len(notifications) > 0 { + notificationCfg = notifications[0] + } + return &BatteryPlugin{ + notifications: notificationCfg, + cfg: cfg, + bus: bus, + logger: logger.With(zap.String("plugin", "battery")), + } +} + +// BatteryBody represents the body of a kdeconnect.battery packet. +// A body carrying only Request asks the peer to report its state; it is +// not a state update and must never touch the stored charge. +type BatteryBody struct { + CurrentCharge int `json:"currentCharge"` + IsCharging bool `json:"isCharging"` + ThresholdEvent int `json:"thresholdEvent"` + Request bool `json:"request,omitempty"` +} + +func (p *BatteryPlugin) Name() string { return "Battery" } +func (p *BatteryPlugin) Timeout() time.Duration { return 5 * time.Second } +func (p *BatteryPlugin) IncomingTypes() []string { + return []string{"kdeconnect.battery", "kdeconnect.battery.request"} +} +func (p *BatteryPlugin) OutgoingTypes() []string { + return []string{"kdeconnect.battery", "kdeconnect.battery.request"} +} diff --git a/internal/plugins/clipboard/backend.go b/internal/plugins/clipboard/backend.go new file mode 100644 index 0000000..aaf936e --- /dev/null +++ b/internal/plugins/clipboard/backend.go @@ -0,0 +1,165 @@ +package clipboard + +import ( + "bytes" + "context" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "strings" + "time" + + "go.uber.org/zap" +) + +// probeBackend determines the usable clipboard backend (wl-paste/xclip) by +// inspecting the environment and $XDG_RUNTIME_DIR. It is side-effect free and +// unit-testable. The returned string is the WAYLAND_DISPLAY value to inject +// into spawned subprocesses (empty for X11/unknown). +// +// A Wayland socket is preferred over DISPLAY: under a systemd user service +// WAYLAND_DISPLAY is often unset at startup, and treating DISPLAY as the +// backend on a Wayland session silently copies to the X clipboard where +// Wayland-native apps never see it. +func probeBackend() (clipboardBackend, string) { + rtDir := os.Getenv("XDG_RUNTIME_DIR") + + // Wayland: trust WAYLAND_DISPLAY only if its socket actually exists, + // otherwise scan the runtime dir for any live wayland-* socket. + if disp := os.Getenv("WAYLAND_DISPLAY"); disp != "" && rtDir != "" { + if _, err := os.Stat(filepath.Join(rtDir, disp)); err == nil { + if _, err := exec.LookPath("wl-paste"); err == nil { + return backendWayland, disp + } + } + } + if rtDir != "" { + if entries, err := os.ReadDir(rtDir); err == nil { + for _, e := range entries { + name := e.Name() + if !e.IsDir() && strings.HasPrefix(name, "wayland-") { + if _, err := exec.LookPath("wl-paste"); err == nil { + return backendWayland, name + } + } + } + } + } + + // X11 fallback. + if os.Getenv("DISPLAY") != "" { + if _, err := exec.LookPath("xclip"); err == nil { + return backendX11, "" + } + } + return backendUnknown, "" +} + +// getBackend returns the cached backend, re-probing while it is unknown so an +// early failed probe (compositor not up yet, env not imported) does not stick +// for the lifetime of the process. Only non-unknown results are cached. +func (p *ClipboardPlugin) getBackend() (clipboardBackend, string) { + p.mu.Lock() + defer p.mu.Unlock() + if p.backend != backendUnknown { + return p.backend, p.wlDisplay + } + backend, disp := p.probe() + if backend != backendUnknown { + p.backend = backend + p.wlDisplay = disp + p.logger.Debug("clipboard: backend detected", + zap.Int("backend", int(backend)), zap.String("wl_display", disp)) + } + return backend, disp +} + +// clipboardCmd builds an exec.Cmd for a clipboard tool with WAYLAND_DISPLAY +// injected into the subprocess environment from the probed socket, so the +// tool works even when the daemon's own environment lacks the variable. +// The command itself carries no deadline; runClipboard applies the timeout +// around execution so a hung wl-paste can never stall the daemon. +func (p *ClipboardPlugin) clipboardCmd(name string, args ...string) *exec.Cmd { + // A background context: the deadline lives in runClipboard, which binds + // a bounded context around execution so a hung tool is killed there. + cmd := exec.CommandContext(context.Background(), name, args...) + cmd.WaitDelay = time.Second + if _, disp := p.getBackend(); disp != "" { + cmd.Env = append(os.Environ(), "WAYLAND_DISPLAY="+disp) + } + return cmd +} + +// runClipboard runs a clipboard subprocess bounded by the clipboard timeout +// and captures any stderr the tool prints. On failure the stderr is wrapped +// into the returned error so the real reason (e.g. "No selection", a compositor +// error) is visible instead of a bare "exit status N". +func (p *ClipboardPlugin) runClipboard(ctx context.Context, cmd *exec.Cmd) ([]byte, error) { + tctx, cancel := context.WithTimeout(ctx, clipboardTimeout) + defer cancel() + + timed := exec.CommandContext(tctx, cmd.Path, cmd.Args[1:]...) + timed.Env = cmd.Env + timed.WaitDelay = capKillDelay(cmd.WaitDelay) + timed.Stdin = nil + + var stderr bytes.Buffer + timed.Stderr = &stderr + + out, err := timed.Output() + if err != nil { + if msg := strings.TrimSpace(stderr.String()); msg != "" { + err = fmt.Errorf("%w: %s", err, msg) + } + } + return out, err +} + +// runCopy writes clipboard data via wl-copy/xclip -i. Unlike runClipboard it +// must NOT capture stdout/stderr through os.Pipe: wl-copy forks a persistent +// background manager that inherits the pipe fds, so the pipe never EOFs and +// cmd.Output()/cmd.Wait() would block until the manager dies. Pointing the +// child's fds at the null device instead makes Run() return as soon as the +// forking wl-copy process exits. wl-paste never forks, which is why the read +// path above can still use pipes. +func (p *ClipboardPlugin) runCopy(ctx context.Context, cmd *exec.Cmd, stdin io.Reader) error { + tctx, cancel := context.WithTimeout(ctx, clipboardTimeout) + defer cancel() + + timed := exec.CommandContext(tctx, cmd.Path, cmd.Args[1:]...) + timed.Env = cmd.Env + timed.WaitDelay = capKillDelay(cmd.WaitDelay) + timed.Stdin = stdin + timed.Stdout = nil + timed.Stderr = nil + return timed.Run() +} + +// capKillDelay returns nonzero bound on how long Wait may block on the +// stdout/stderr pipes after the timeout kills the child. Callers built via +// clipboardCmd start with WaitDelay=1s, but a bare exec.CommandContext +// (tests, or any future caller) defaults to 0, which means Wait blocks until +// the pipes EOF — and a killed child that forked a grandchild holding those +// fds would pin runClipboard open until the grandchild exits (e.g. a whole +// 30s sleep). Ensure it always returns promptly after the deadline. +func capKillDelay(d time.Duration) time.Duration { + if d <= 0 { + return time.Second + } + return d +} + +// isNoSelection reports whether a clipboard tool failure actually means the +// clipboard is empty (nothing to push) rather than a real problem with the +// tool or compositor. Matches wl-paste's and xclip's "nothing here" messages. +func isNoSelection(err error) bool { + if err == nil { + return false + } + s := strings.ToLower(err.Error()) + return strings.Contains(s, "no selection") || + strings.Contains(s, "nothing is copied") || + strings.Contains(s, "no data") +} diff --git a/internal/plugins/clipboard/clipboard.go b/internal/plugins/clipboard/clipboard.go deleted file mode 100644 index b894f57..0000000 --- a/internal/plugins/clipboard/clipboard.go +++ /dev/null @@ -1,509 +0,0 @@ -package clipboard - -import ( - "bytes" - "context" - "crypto/tls" - "encoding/json" - "fmt" - "io" - "mime" - "net" - "os" - "os/exec" - "path/filepath" - "strings" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/protocol" - "github.com/bethropolis/kcd/internal/transport" - "go.uber.org/zap" -) - -type clipboardBackend int - -const ( - backendUnknown clipboardBackend = iota - backendWayland - backendX11 -) - -// ClipboardPlugin handles clipboard sync both directions. -type ClipboardPlugin struct { - sidechannel transport.SidechannelOptions - pushOnConnect bool - lastTimestamp int64 - tlsConfig *tls.Config - logger *zap.Logger - backend clipboardBackend - wlDisplay string // WAYLAND_DISPLAY value for spawned subprocesses - probe func() (clipboardBackend, string) - mu sync.Mutex - lastContent string // last content received from phone (inbound) - lastPushedContent string // last content sent to phone (outbound) -} - -// NewClipboardPlugin creates a clipboard plugin. -func NewClipboardPlugin(tlsConfig *tls.Config, logger *zap.Logger, pushOnConnect bool, options ...transport.SidechannelOptions) *ClipboardPlugin { - var sidechannel transport.SidechannelOptions - if len(options) > 0 { - sidechannel = options[0] - } - if logger == nil { - logger = zap.NewNop() - } - return &ClipboardPlugin{ - sidechannel: sidechannel, - tlsConfig: tlsConfig, - pushOnConnect: pushOnConnect, - logger: logger.With(zap.String("plugin", "clipboard")), - probe: probeBackend, - } -} - -// probeBackend determines the usable clipboard backend (wl-paste/xclip) by -// inspecting the environment and $XDG_RUNTIME_DIR. It is side-effect free and -// unit-testable. The returned string is the WAYLAND_DISPLAY value to inject -// into spawned subprocesses (empty for X11/unknown). -// -// A Wayland socket is preferred over DISPLAY: under a systemd user service -// WAYLAND_DISPLAY is often unset at startup, and treating DISPLAY as the -// backend on a Wayland session silently copies to the X clipboard where -// Wayland-native apps never see it. -func probeBackend() (clipboardBackend, string) { - rtDir := os.Getenv("XDG_RUNTIME_DIR") - - // Wayland: trust WAYLAND_DISPLAY only if its socket actually exists, - // otherwise scan the runtime dir for any live wayland-* socket. - if disp := os.Getenv("WAYLAND_DISPLAY"); disp != "" && rtDir != "" { - if _, err := os.Stat(filepath.Join(rtDir, disp)); err == nil { - if _, err := exec.LookPath("wl-paste"); err == nil { - return backendWayland, disp - } - } - } - if rtDir != "" { - if entries, err := os.ReadDir(rtDir); err == nil { - for _, e := range entries { - name := e.Name() - if !e.IsDir() && strings.HasPrefix(name, "wayland-") { - if _, err := exec.LookPath("wl-paste"); err == nil { - return backendWayland, name - } - } - } - } - } - - // X11 fallback. - if os.Getenv("DISPLAY") != "" { - if _, err := exec.LookPath("xclip"); err == nil { - return backendX11, "" - } - } - return backendUnknown, "" -} - -// getBackend returns the cached backend, re-probing while it is unknown so an -// early failed probe (compositor not up yet, env not imported) does not stick -// for the lifetime of the process. Only non-unknown results are cached. -func (p *ClipboardPlugin) getBackend() (clipboardBackend, string) { - p.mu.Lock() - defer p.mu.Unlock() - if p.backend != backendUnknown { - return p.backend, p.wlDisplay - } - backend, disp := p.probe() - if backend != backendUnknown { - p.backend = backend - p.wlDisplay = disp - p.logger.Debug("clipboard: backend detected", - zap.Int("backend", int(backend)), zap.String("wl_display", disp)) - } - return backend, disp -} - -// clipboardCmd builds an exec.Cmd for a clipboard tool with WAYLAND_DISPLAY -// injected into the subprocess environment from the probed socket, so the -// tool works even when the daemon's own environment lacks the variable. -// The command itself carries no deadline; runClipboard applies the timeout -// around execution so a hung wl-paste can never stall the daemon. -func (p *ClipboardPlugin) clipboardCmd(name string, args ...string) *exec.Cmd { - // A background context: the deadline lives in runClipboard, which binds - // a bounded context around execution so a hung tool is killed there. - cmd := exec.CommandContext(context.Background(), name, args...) - cmd.WaitDelay = time.Second - if _, disp := p.getBackend(); disp != "" { - cmd.Env = append(os.Environ(), "WAYLAND_DISPLAY="+disp) - } - return cmd -} - -// runClipboard runs a clipboard subprocess bounded by the clipboard timeout -// and captures any stderr the tool prints. On failure the stderr is wrapped -// into the returned error so the real reason (e.g. "No selection", a compositor -// error) is visible instead of a bare "exit status N". -func (p *ClipboardPlugin) runClipboard(ctx context.Context, cmd *exec.Cmd) ([]byte, error) { - tctx, cancel := context.WithTimeout(ctx, clipboardTimeout) - defer cancel() - - timed := exec.CommandContext(tctx, cmd.Path, cmd.Args[1:]...) - timed.Env = cmd.Env - timed.WaitDelay = capKillDelay(cmd.WaitDelay) - timed.Stdin = nil - - var stderr bytes.Buffer - timed.Stderr = &stderr - - out, err := timed.Output() - if err != nil { - if msg := strings.TrimSpace(stderr.String()); msg != "" { - err = fmt.Errorf("%w: %s", err, msg) - } - } - return out, err -} - -// runCopy writes clipboard data via wl-copy/xclip -i. Unlike runClipboard it -// must NOT capture stdout/stderr through os.Pipe: wl-copy forks a persistent -// background manager that inherits the pipe fds, so the pipe never EOFs and -// cmd.Output()/cmd.Wait() would block until the manager dies. Pointing the -// child's fds at the null device instead makes Run() return as soon as the -// forking wl-copy process exits. wl-paste never forks, which is why the read -// path above can still use pipes. -func (p *ClipboardPlugin) runCopy(ctx context.Context, cmd *exec.Cmd, stdin io.Reader) error { - tctx, cancel := context.WithTimeout(ctx, clipboardTimeout) - defer cancel() - - timed := exec.CommandContext(tctx, cmd.Path, cmd.Args[1:]...) - timed.Env = cmd.Env - timed.WaitDelay = capKillDelay(cmd.WaitDelay) - timed.Stdin = stdin - timed.Stdout = nil - timed.Stderr = nil - return timed.Run() -} - -// capKillDelay returns nonzero bound on how long Wait may block on the -// stdout/stderr pipes after the timeout kills the child. Callers built via -// clipboardCmd start with WaitDelay=1s, but a bare exec.CommandContext -// (tests, or any future caller) defaults to 0, which means Wait blocks until -// the pipes EOF — and a killed child that forked a grandchild holding those -// fds would pin runClipboard open until the grandchild exits (e.g. a whole -// 30s sleep). Ensure it always returns promptly after the deadline. -func capKillDelay(d time.Duration) time.Duration { - if d <= 0 { - return time.Second - } - return d -} - -// isNoSelection reports whether a clipboard tool failure actually means the -// clipboard is empty (nothing to push) rather than a real problem with the -// tool or compositor. Matches wl-paste's and xclip's "nothing here" messages. -func isNoSelection(err error) bool { - if err == nil { - return false - } - s := strings.ToLower(err.Error()) - return strings.Contains(s, "no selection") || - strings.Contains(s, "nothing is copied") || - strings.Contains(s, "no data") -} - -// clipboardTimeout bounds every wl-copy/wl-paste/xclip subprocess so a hung -// clipboard tool cannot block clipboard sync indefinitely. -const clipboardTimeout = 2 * time.Second - -// ClipboardBody represents the content of a clipboard packet. -type ClipboardBody struct { - Content string `json:"content"` - Timestamp int64 `json:"timestamp,omitempty"` -} - -// Name returns the plugin name. -func (p *ClipboardPlugin) Name() string { return "Clipboard" } - -// Timeout returns the timeout. -func (p *ClipboardPlugin) Timeout() time.Duration { return 5 * time.Second } - -// IncomingTypes returns the packet types this plugin handles. -func (p *ClipboardPlugin) IncomingTypes() []string { - return []string{"kdeconnect.clipboard", "kdeconnect.clipboard.connect", "kdeconnect.clipboard.file"} -} - -// OutgoingTypes returns the packet types this plugin may send. -func (p *ClipboardPlugin) OutgoingTypes() []string { - return []string{"kdeconnect.clipboard"} -} - -// Handle processes incoming clipboard packets. -func (p *ClipboardPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - // Handle image/file clipboard transfer - if pkt.Type == "kdeconnect.clipboard.file" { - return p.handleClipboardFile(ctx, dev, pkt) - } - - var body ClipboardBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - // Ignore if it's just connectivity notification (kdeconnect.clipboard.connect) - // but failed to parse content. - return nil - } - - if body.Content == "" { - return nil - } - - p.mu.Lock() - if body.Content == p.lastContent { - p.mu.Unlock() - return nil - } - // Guard lastTimestamp under the same lock as lastContent — they form a - // consistent pair and both fields are written from the TCP read goroutine. - if body.Timestamp > 0 { - if body.Timestamp < p.lastTimestamp { - p.mu.Unlock() - return nil - } - p.lastTimestamp = body.Timestamp - } - p.lastContent = body.Content - p.mu.Unlock() - - // Spawning goroutine as Handlers must not block. - go func() { - switch backend, _ := p.getBackend(); backend { - case backendWayland: - // -n: wl-copy appends a trailing newline by default. Without it the - // local selection becomes content+"\n", which differs from the - // inbound lastContent guard and makes --watch echo the phone's own - // clipboard straight back to it. - if err := p.runCopy(context.Background(), p.clipboardCmd("wl-copy", "-n"), strings.NewReader(body.Content)); err != nil { - p.logger.Warn("clipboard: failed to set clipboard", zap.Error(err)) - } - case backendX11: - if err := p.runCopy(context.Background(), p.clipboardCmd("xclip", "-selection", "clipboard"), strings.NewReader(body.Content)); err != nil { - p.logger.Warn("clipboard: failed to set clipboard", zap.Error(err)) - } - default: - p.logger.Debug("clipboard: no backend available, dropping inbound copy") - } - }() - - return nil -} - -// ClipboardFileBody is the body of kdeconnect.clipboard.file. -type ClipboardFileBody struct { - Filename string `json:"filename"` -} - -const maxClipboardFileSize = 50 * 1024 * 1024 // 50 MB safety limit (matches KDE Connect C++) - -func (p *ClipboardPlugin) handleClipboardFile(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - if pkt.PayloadSize <= 0 || pkt.PayloadTransferInfo == nil { - return nil - } - - if pkt.PayloadSize > maxClipboardFileSize { - p.logger.Warn("clipboard file: rejected payload exceeding size limit", - zap.Int64("size", pkt.PayloadSize), - zap.Int("limit_bytes", maxClipboardFileSize), - ) - return nil - } - var body ClipboardFileBody - if err := json.Unmarshal(pkt.Body, &body); err != nil || body.Filename == "" { - return nil - } - - remoteIP := dev.RemoteIP() - if remoteIP == nil { - return nil - } - - payloadSize := pkt.PayloadSize - payloadPort := pkt.PayloadTransferInfo.Port - filename := body.Filename - expectedFP := cert.PinnedFingerprint(dev.PeerCert()) - - go func() { - // Download to a temp file - tmpFile, err := os.CreateTemp("", "kcd-clip-*"+filepath.Ext(filename)) - if err != nil { - p.logger.Error("clipboard file: failed to create temp file", zap.Error(err)) - return - } - tmpPath := tmpFile.Name() - tmpFile.Close() - defer os.Remove(tmpPath) - - if err := downloadToFile(ctx, remoteIP, payloadPort, payloadSize, tmpPath, p.tlsConfig, expectedFP, p.logger, p.sidechannel); err != nil { - p.logger.Error("clipboard file: download failed", zap.Error(err)) - return - } - - // Detect MIME type from extension - mimeType := mime.TypeByExtension(filepath.Ext(filename)) - if mimeType == "" { - mimeType = "application/octet-stream" - } - - t, err := os.Open(tmpPath) - if err != nil { - return - } - defer t.Close() - - var cmd *exec.Cmd - switch backend, _ := p.getBackend(); backend { - case backendWayland: - cmd = p.clipboardCmd("wl-copy", "-n", "--type", mimeType) - case backendX11: - cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-t", mimeType, "-i") - default: - return - } - if err := p.runCopy(context.Background(), cmd, t); err != nil { - p.logger.Warn("clipboard file: failed to set clipboard", zap.Error(err)) - } - }() - - return nil -} - -// downloadToFile dials a TLS side-channel and streams the payload to dest. -func downloadToFile(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger, options ...transport.SidechannelOptions) error { - conn, err := transport.DialSidechannel(ctx, ip, port, tlsConfig, expectedFP, logger, options...) - if err != nil { - return err - } - defer conn.Close() - - f, err := os.Create(dest) - if err != nil { - return fmt.Errorf("clipboard: create file %s: %w", dest, err) - } - defer f.Close() - - _, err = io.Copy(f, io.LimitReader(conn, size)) - if err != nil { - os.Remove(dest) // don't leave a corrupt partial behind - return fmt.Errorf("clipboard: stream to %s: %w", dest, err) - } - return nil -} - -// Push copies the local clipboard to the remote device using wl-paste or xclip -o. -func Push(ctx context.Context, dev device.Sender, p *ClipboardPlugin) error { - var cmd *exec.Cmd - switch backend, _ := p.getBackend(); backend { - case backendWayland: - cmd = p.clipboardCmd("wl-paste", "-n") - case backendX11: - cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-o") - default: - return fmt.Errorf("clipboard: no clipboard tool available") - } - - out, err := p.runClipboard(ctx, cmd) - if err != nil { - // An empty clipboard (fresh session, nothing copied yet) is not an - // error worth failing a push for — the tool exits non-zero with a - // "no selection"-style message. Treat it as nothing to push so the - // CLI and --watch don't spam errors until something is copied. - if isNoSelection(err) { - return nil - } - return err - } - - content := string(out) - if content == "" { - return nil - } - - p.mu.Lock() - // Skip if content matches what we last received from the phone (lastContent) - // OR what we last pushed outbound (lastPushedContent). - // - // lastContent guard: prevents sending the phone's own content back. - // lastPushedContent guard: prevents duplicate pushes when the local - // clipboard hasn't changed between two Push calls. - // - // Comparisons are normalized of trailing newlines so a stray \n appended - // by some clipboard tooling cannot silently disable the guard and cause - // an echo of the phone's own content back to it. - if normClip(content) == normClip(p.lastContent) || normClip(content) == normClip(p.lastPushedContent) { - p.mu.Unlock() - return nil - } - p.lastPushedContent = content - p.mu.Unlock() - - pkt, err := protocol.NewPacket("kdeconnect.clipboard", ClipboardBody{ - Content: content, - }) - if err != nil { - return err - } - - // All outgoing packets must use device.Send. - return dev.Send(pkt) -} - -// normClip strips trailing newlines so content read back from wl-copy/xclip -// compares equal to the raw content we stored, regardless of which tool -// appended (or not) a trailing newline. -func normClip(s string) string { - return strings.TrimRight(s, "\n") -} - -func (p *ClipboardPlugin) readClipboard() string { - var cmd *exec.Cmd - switch backend, _ := p.getBackend(); backend { - case backendWayland: - cmd = p.clipboardCmd("wl-paste", "-n") - case backendX11: - cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-o") - default: - return "" - } - out, err := p.runClipboard(context.Background(), cmd) - if err != nil { - p.logger.Debug("clipboard: read failed", zap.Error(err)) - return "" - } - return string(out) -} - -func (p *ClipboardPlugin) OnConnect(dev device.Sender) { - if !p.pushOnConnect { - return - } - content := p.readClipboard() - - p.mu.Lock() - p.lastPushedContent = content - p.mu.Unlock() - - body := ClipboardBody{ - Content: content, - Timestamp: time.Now().UnixMilli(), - } - pkt, err := protocol.NewPacket("kdeconnect.clipboard.connect", body) - if err != nil { - p.logger.Debug("clipboard: OnConnect: failed to build packet", zap.Error(err)) - return - } - // Best-effort — device may still be completing the TLS handshake. - _ = dev.Send(pkt) -} - -func (p *ClipboardPlugin) OnDisconnect(dev device.Sender) { -} diff --git a/internal/plugins/clipboard/handle.go b/internal/plugins/clipboard/handle.go new file mode 100644 index 0000000..90bf11a --- /dev/null +++ b/internal/plugins/clipboard/handle.go @@ -0,0 +1,180 @@ +package clipboard + +import ( + "context" + "crypto/tls" + "encoding/json" + "fmt" + "io" + "mime" + "net" + "os" + "os/exec" + "path/filepath" + "strings" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// Handle processes incoming clipboard packets. +func (p *ClipboardPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + // Handle image/file clipboard transfer + if pkt.Type == "kdeconnect.clipboard.file" { + return p.handleClipboardFile(ctx, dev, pkt) + } + + var body ClipboardBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + // Ignore if it's just connectivity notification (kdeconnect.clipboard.connect) + // but failed to parse content. + return nil + } + + if body.Content == "" { + return nil + } + + p.mu.Lock() + if body.Content == p.lastContent { + p.mu.Unlock() + return nil + } + // Guard lastTimestamp under the same lock as lastContent — they form a + // consistent pair and both fields are written from the TCP read goroutine. + if body.Timestamp > 0 { + if body.Timestamp < p.lastTimestamp { + p.mu.Unlock() + return nil + } + p.lastTimestamp = body.Timestamp + } + p.lastContent = body.Content + p.mu.Unlock() + + // Spawning goroutine as Handlers must not block. + go func() { + switch backend, _ := p.getBackend(); backend { + case backendWayland: + // -n: wl-copy appends a trailing newline by default. Without it the + // local selection becomes content+"\n", which differs from the + // inbound lastContent guard and makes --watch echo the phone's own + // clipboard straight back to it. + if err := p.runCopy(context.Background(), p.clipboardCmd("wl-copy", "-n"), strings.NewReader(body.Content)); err != nil { + p.logger.Warn("clipboard: failed to set clipboard", zap.Error(err)) + } + case backendX11: + if err := p.runCopy(context.Background(), p.clipboardCmd("xclip", "-selection", "clipboard"), strings.NewReader(body.Content)); err != nil { + p.logger.Warn("clipboard: failed to set clipboard", zap.Error(err)) + } + default: + p.logger.Debug("clipboard: no backend available, dropping inbound copy") + } + }() + + return nil +} + +// ClipboardFileBody is the body of kdeconnect.clipboard.file. +type ClipboardFileBody struct { + Filename string `json:"filename"` +} + +const maxClipboardFileSize = 50 * 1024 * 1024 // 50 MB safety limit (matches KDE Connect C++) + +func (p *ClipboardPlugin) handleClipboardFile(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + if pkt.PayloadSize <= 0 || pkt.PayloadTransferInfo == nil { + return nil + } + + if pkt.PayloadSize > maxClipboardFileSize { + p.logger.Warn("clipboard file: rejected payload exceeding size limit", + zap.Int64("size", pkt.PayloadSize), + zap.Int("limit_bytes", maxClipboardFileSize), + ) + return nil + } + var body ClipboardFileBody + if err := json.Unmarshal(pkt.Body, &body); err != nil || body.Filename == "" { + return nil + } + + remoteIP := dev.RemoteIP() + if remoteIP == nil { + return nil + } + + payloadSize := pkt.PayloadSize + payloadPort := pkt.PayloadTransferInfo.Port + filename := body.Filename + expectedFP := cert.PinnedFingerprint(dev.PeerCert()) + + go func() { + // Download to a temp file + tmpFile, err := os.CreateTemp("", "kcd-clip-*"+filepath.Ext(filename)) + if err != nil { + p.logger.Error("clipboard file: failed to create temp file", zap.Error(err)) + return + } + tmpPath := tmpFile.Name() + tmpFile.Close() + defer os.Remove(tmpPath) + + if err := downloadToFile(ctx, remoteIP, payloadPort, payloadSize, tmpPath, p.tlsConfig, expectedFP, p.logger, p.sidechannel); err != nil { + p.logger.Error("clipboard file: download failed", zap.Error(err)) + return + } + + // Detect MIME type from extension + mimeType := mime.TypeByExtension(filepath.Ext(filename)) + if mimeType == "" { + mimeType = "application/octet-stream" + } + + t, err := os.Open(tmpPath) + if err != nil { + return + } + defer t.Close() + + var cmd *exec.Cmd + switch backend, _ := p.getBackend(); backend { + case backendWayland: + cmd = p.clipboardCmd("wl-copy", "-n", "--type", mimeType) + case backendX11: + cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-t", mimeType, "-i") + default: + return + } + if err := p.runCopy(context.Background(), cmd, t); err != nil { + p.logger.Warn("clipboard file: failed to set clipboard", zap.Error(err)) + } + }() + + return nil +} + +// downloadToFile dials a TLS side-channel and streams the payload to dest. +func downloadToFile(ctx context.Context, ip net.IP, port int, size int64, dest string, tlsConfig *tls.Config, expectedFP string, logger *zap.Logger, options ...transport.SidechannelOptions) error { + conn, err := transport.DialSidechannel(ctx, ip, port, tlsConfig, expectedFP, logger, options...) + if err != nil { + return err + } + defer conn.Close() + + f, err := os.Create(dest) + if err != nil { + return fmt.Errorf("clipboard: create file %s: %w", dest, err) + } + defer f.Close() + + _, err = io.Copy(f, io.LimitReader(conn, size)) + if err != nil { + os.Remove(dest) // don't leave a corrupt partial behind + return fmt.Errorf("clipboard: stream to %s: %w", dest, err) + } + return nil +} diff --git a/internal/plugins/clipboard/push.go b/internal/plugins/clipboard/push.go new file mode 100644 index 0000000..f292f57 --- /dev/null +++ b/internal/plugins/clipboard/push.go @@ -0,0 +1,122 @@ +package clipboard + +import ( + "context" + "fmt" + "os/exec" + "strings" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// Push copies the local clipboard to the remote device using wl-paste or xclip -o. +func Push(ctx context.Context, dev device.Sender, p *ClipboardPlugin) error { + var cmd *exec.Cmd + switch backend, _ := p.getBackend(); backend { + case backendWayland: + cmd = p.clipboardCmd("wl-paste", "-n") + case backendX11: + cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-o") + default: + return fmt.Errorf("clipboard: no clipboard tool available") + } + + out, err := p.runClipboard(ctx, cmd) + if err != nil { + // An empty clipboard (fresh session, nothing copied yet) is not an + // error worth failing a push for — the tool exits non-zero with a + // "no selection"-style message. Treat it as nothing to push so the + // CLI and --watch don't spam errors until something is copied. + if isNoSelection(err) { + return nil + } + return err + } + + content := string(out) + if content == "" { + return nil + } + + p.mu.Lock() + // Skip if content matches what we last received from the phone (lastContent) + // OR what we last pushed outbound (lastPushedContent). + // + // lastContent guard: prevents sending the phone's own content back. + // lastPushedContent guard: prevents duplicate pushes when the local + // clipboard hasn't changed between two Push calls. + // + // Comparisons are normalized of trailing newlines so a stray \n appended + // by some clipboard tooling cannot silently disable the guard and cause + // an echo of the phone's own content back to it. + if normClip(content) == normClip(p.lastContent) || normClip(content) == normClip(p.lastPushedContent) { + p.mu.Unlock() + return nil + } + p.lastPushedContent = content + p.mu.Unlock() + + pkt, err := protocol.NewPacket("kdeconnect.clipboard", ClipboardBody{ + Content: content, + }) + if err != nil { + return err + } + + // All outgoing packets must use device.Send. + return dev.Send(pkt) +} + +// normClip strips trailing newlines so content read back from wl-copy/xclip +// compares equal to the raw content we stored, regardless of which tool +// appended (or not) a trailing newline. +func normClip(s string) string { + return strings.TrimRight(s, "\n") +} + +func (p *ClipboardPlugin) readClipboard() string { + var cmd *exec.Cmd + switch backend, _ := p.getBackend(); backend { + case backendWayland: + cmd = p.clipboardCmd("wl-paste", "-n") + case backendX11: + cmd = p.clipboardCmd("xclip", "-selection", "clipboard", "-o") + default: + return "" + } + out, err := p.runClipboard(context.Background(), cmd) + if err != nil { + p.logger.Debug("clipboard: read failed", zap.Error(err)) + return "" + } + return string(out) +} + +func (p *ClipboardPlugin) OnConnect(dev device.Sender) { + if !p.pushOnConnect { + return + } + content := p.readClipboard() + + p.mu.Lock() + p.lastPushedContent = content + p.mu.Unlock() + + body := ClipboardBody{ + Content: content, + Timestamp: time.Now().UnixMilli(), + } + pkt, err := protocol.NewPacket("kdeconnect.clipboard.connect", body) + if err != nil { + p.logger.Debug("clipboard: OnConnect: failed to build packet", zap.Error(err)) + return + } + // Best-effort — device may still be completing the TLS handshake. + _ = dev.Send(pkt) +} + +func (p *ClipboardPlugin) OnDisconnect(dev device.Sender) { +} diff --git a/internal/plugins/clipboard/types.go b/internal/plugins/clipboard/types.go new file mode 100644 index 0000000..a259efa --- /dev/null +++ b/internal/plugins/clipboard/types.go @@ -0,0 +1,77 @@ +package clipboard + +import ( + "crypto/tls" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +type clipboardBackend int + +const ( + backendUnknown clipboardBackend = iota + backendWayland + backendX11 +) + +// ClipboardPlugin handles clipboard sync both directions. +type ClipboardPlugin struct { + sidechannel transport.SidechannelOptions + pushOnConnect bool + lastTimestamp int64 + tlsConfig *tls.Config + logger *zap.Logger + backend clipboardBackend + wlDisplay string // WAYLAND_DISPLAY value for spawned subprocesses + probe func() (clipboardBackend, string) + mu sync.Mutex + lastContent string // last content received from phone (inbound) + lastPushedContent string // last content sent to phone (outbound) +} + +// NewClipboardPlugin creates a clipboard plugin. +func NewClipboardPlugin(tlsConfig *tls.Config, logger *zap.Logger, pushOnConnect bool, options ...transport.SidechannelOptions) *ClipboardPlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } + if logger == nil { + logger = zap.NewNop() + } + return &ClipboardPlugin{ + sidechannel: sidechannel, + tlsConfig: tlsConfig, + pushOnConnect: pushOnConnect, + logger: logger.With(zap.String("plugin", "clipboard")), + probe: probeBackend, + } +} + +// clipboardTimeout bounds every wl-copy/wl-paste/xclip subprocess so a hung +// clipboard tool cannot block clipboard sync indefinitely. +const clipboardTimeout = 2 * time.Second + +// ClipboardBody represents the content of a clipboard packet. +type ClipboardBody struct { + Content string `json:"content"` + Timestamp int64 `json:"timestamp,omitempty"` +} + +// Name returns the plugin name. +func (p *ClipboardPlugin) Name() string { return "Clipboard" } + +// Timeout returns the timeout. +func (p *ClipboardPlugin) Timeout() time.Duration { return 5 * time.Second } + +// IncomingTypes returns the packet types this plugin handles. +func (p *ClipboardPlugin) IncomingTypes() []string { + return []string{"kdeconnect.clipboard", "kdeconnect.clipboard.connect", "kdeconnect.clipboard.file"} +} + +// OutgoingTypes returns the packet types this plugin may send. +func (p *ClipboardPlugin) OutgoingTypes() []string { + return []string{"kdeconnect.clipboard"} +} diff --git a/internal/plugins/contacts/contacts.go b/internal/plugins/contacts/contacts.go deleted file mode 100644 index 5934ccd..0000000 --- a/internal/plugins/contacts/contacts.go +++ /dev/null @@ -1,594 +0,0 @@ -package contacts - -import ( - "context" - "encoding/json" - "fmt" - "os" - "path/filepath" - "regexp" - "sort" - "strconv" - "strings" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "go.uber.org/zap" -) - -// KDE Connect contacts packet types (see kdeconnect-kde -// plugins/contacts/contactsplugin.h and the Android ContactsPlugin). -const ( - PacketTypeContactsRequestUIDs = "kdeconnect.contacts.request_all_uids_timestamps" - PacketTypeContactsRequestVCards = "kdeconnect.contacts.request_vcards_by_uid" - PacketTypeContactsResponseUIDs = "kdeconnect.contacts.response_uids_timestamps" - PacketTypeContactsResponseVCards = "kdeconnect.contacts.response_vcards" -) - -const ( - // maxContactUIDs bounds a single uids response (address books are in - // the hundreds; this is two orders of magnitude of headroom). - maxContactUIDs = 20000 - // maxVCardBytes caps one vCard (text contact cards are ~1KB). - maxVCardBytes = 64 * 1024 - // maxSyncBytes caps a whole vcards response before anything hits disk. - maxSyncBytes = 64 * 1024 * 1024 - // vcardChunkSize bounds our outbound request packets. - vcardChunkSize = 100 - // maxDisplayLen truncates contact fields shown to clients. - maxDisplayLen = 256 -) - -// uidSafeChars restricts phone-provided contact UIDs (Android LOOKUP_KEYs) -// to filename-safe characters so they can't escape the cache directory. -var uidSafeChars = regexp.MustCompile(`[^a-zA-Z0-9._-]`) - -// sanitizeUID strips everything but alphanumerics, dot, underscore and -// hyphen. Empty results are rejected by the caller. -func sanitizeUID(uid string) string { - return uidSafeChars.ReplaceAllString(uid, "_") -} - -// ContactSummary is the parsed, display-safe view of one contact. -type ContactSummary struct { - UID string `json:"uid"` - Name string `json:"name"` - Phones []string `json:"phones,omitempty"` - Emails []string `json:"emails,omitempty"` - Timestamp int64 `json:"timestamp"` -} - -// indexEntry is the persisted per-contact record (index.json sidecar, so -// listing never reparses hundreds of vCards). Fetched marks that the vCard -// round completed; entries with only a timestamp are pending fetch. -type indexEntry struct { - Timestamp int64 `json:"timestamp"` - Fetched bool `json:"fetched"` - Name string `json:"name"` - Phones []string `json:"phones,omitempty"` - Emails []string `json:"emails,omitempty"` -} - -// ContactsPlugin syncs the phone address book: it requests UID/timestamp -// lists, fetches vCards for new or changed contacts, caches them per -// device, and deletes stale entries. -type ContactsPlugin struct { - bus *events.Bus - logger *zap.Logger - baseDir string - - // mu serializes sync processing across devices. Syncs are rare; - // one lock avoids per-device lock bookkeeping. - mu sync.Mutex -} - -// NewContactsPlugin creates a contacts plugin caching under -// $XDG_DATA_HOME/kcd/contacts (0600 files, 0700 dirs). -func NewContactsPlugin(bus *events.Bus, logger *zap.Logger, cacheDirs ...string) *ContactsPlugin { - dataHome := os.Getenv("XDG_DATA_HOME") - if dataHome == "" { - home, _ := os.UserHomeDir() - dataHome = filepath.Join(home, ".local", "share") - } - baseDir := filepath.Join(dataHome, "kcd", "contacts") - if len(cacheDirs) > 0 && cacheDirs[0] != "" { - baseDir = cacheDirs[0] - } - _ = os.MkdirAll(baseDir, 0700) - - return &ContactsPlugin{ - bus: bus, - logger: logger.With(zap.String("plugin", "contacts")), - baseDir: baseDir, - } -} - -func (p *ContactsPlugin) Name() string { return "Contacts" } - -func (p *ContactsPlugin) Timeout() time.Duration { return 5 * time.Second } - -func (p *ContactsPlugin) IncomingTypes() []string { - return []string{PacketTypeContactsResponseUIDs, PacketTypeContactsResponseVCards} -} - -func (p *ContactsPlugin) OutgoingTypes() []string { - return []string{PacketTypeContactsRequestUIDs, PacketTypeContactsRequestVCards} -} - -// OnConnect starts a sync for freshly connected paired devices, mirroring -// upstream's connected() -> synchronizeRemoteWithLocal(). -func (p *ContactsPlugin) OnConnect(dev device.Sender) { - if dev.State() != device.StatePaired { - return - } - if err := p.RequestSync(dev); err != nil { - p.logger.Debug("contacts: initial sync request failed", zap.Error(err)) - } -} - -func (p *ContactsPlugin) OnDisconnect(dev device.Sender) {} - -// Handle routes response packets; all parsing and disk I/O runs in a -// worker goroutine so Handle returns immediately (rule 9). -func (p *ContactsPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - switch pkt.Type { - case PacketTypeContactsResponseUIDs: - body := append([]byte(nil), pkt.Body...) - go p.handleUIDsResponse(dev, body) - return nil - case PacketTypeContactsResponseVCards: - body := append([]byte(nil), pkt.Body...) - devID := dev.ID() - go p.handleVCardsResponse(devID, body) - return nil - default: - return nil - } -} - -// RequestSync asks the phone for all contact UIDs and timestamps, which -// starts the sync round trips. Responses arrive async via Handle. -// -// The body must be an empty object, not null: stock implementations send -// `"body":{}` for bodyless requests, and at least one phone build aborts -// the whole link on an explicit null body (observed as an immediate RST -// after every connect that carried `"body":null`). -func (p *ContactsPlugin) RequestSync(dev device.Sender) error { - pkt, err := protocol.NewPacket(PacketTypeContactsRequestUIDs, map[string]any{}) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// requestVCards asks for vCards of the given UIDs, chunked to bound -// outbound packet size. -func (p *ContactsPlugin) requestVCards(dev device.Sender, uids []string) error { - for _, chunk := range chunkUIDs(uids, vcardChunkSize) { - body := map[string]any{"uids": chunk} - pkt, err := protocol.NewPacket(PacketTypeContactsRequestVCards, body) - if err != nil { - return err - } - if err := dev.Send(pkt); err != nil { - return err - } - } - return nil -} - -func chunkUIDs(uids []string, size int) [][]string { - var chunks [][]string - for len(uids) > 0 { - n := size - if len(uids) < n { - n = len(uids) - } - chunks = append(chunks, uids[:n]) - uids = uids[n:] - } - return chunks -} - -// cacheDir resolves the cache dir for a device without creating it. -// The device ID comes from our own registry, but confine defensively. -func (p *ContactsPlugin) cacheDir(deviceID string) (string, error) { - safe := sanitizeUID(deviceID) - if safe == "" { - return "", fmt.Errorf("contacts: unusable device id") - } - dir := filepath.Join(p.baseDir, safe) - if rel, err := filepath.Rel(p.baseDir, dir); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { - return "", fmt.Errorf("contacts: device dir escapes cache") - } - return dir, nil -} - -// deviceDir returns the cache dir for a device, creating it (0700). -func (p *ContactsPlugin) deviceDir(deviceID string) (string, error) { - dir, err := p.cacheDir(deviceID) - if err != nil { - return "", err - } - if err := os.MkdirAll(dir, 0700); err != nil { - return "", fmt.Errorf("contacts: create dir: %w", err) - } - return dir, nil -} - -// contactPath resolves the vcf file for a UID inside dir, confined to dir. -func contactPath(dir, uid string) (string, error) { - safe := sanitizeUID(uid) - if safe == "" { - return "", fmt.Errorf("contacts: unusable uid") - } - path := filepath.Join(dir, safe+".vcf") - if rel, err := filepath.Rel(dir, path); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { - return "", fmt.Errorf("contacts: uid escapes cache") - } - return path, nil -} - -// loadIndex reads the sidecar index (missing file = empty cache, not error). -func loadIndex(dir string) map[string]indexEntry { - idx := make(map[string]indexEntry) - data, err := os.ReadFile(filepath.Join(dir, "index.json")) - if err != nil { - return idx - } - _ = json.Unmarshal(data, &idx) - if idx == nil { - idx = make(map[string]indexEntry) - } - return idx -} - -// saveIndex persists the sidecar index (0600). -func saveIndex(dir string, idx map[string]indexEntry) error { - data, err := json.MarshalIndent(idx, "", " ") - if err != nil { - return err - } - return os.WriteFile(filepath.Join(dir, "index.json"), data, 0600) -} - -// coerceTimestamp parses Android's string-encoded timestamps as well as -// plain JSON numbers (the header doc shows ints). Garbage yields 0, which -// forces a re-fetch — the safe direction, never a false "unchanged". -func coerceTimestamp(raw json.RawMessage) int64 { - var asAny any - if err := json.Unmarshal(raw, &asAny); err != nil { - return 0 - } - switch v := asAny.(type) { - case string: - n, err := strconv.ParseInt(strings.TrimSpace(v), 10, 64) - if err != nil { - return 0 - } - return n - case float64: - if v < 0 { - return 0 - } - return int64(v) - default: - return 0 - } -} - -// handleUIDsResponse diffs the phone's UID/timestamp list against the local -// cache, deletes stale entries, and requests vCards for new/changed ones. -// dev is the live sender (device.Send is channel-safe across goroutines). -func (p *ContactsPlugin) handleUIDsResponse(dev device.Sender, body []byte) { - deviceID := dev.ID() - - // Cache diff under lock; the vCard request round below does network - // I/O and must not hold it. - toFetch, added, updated, deleted := func() ([]string, int, int, int) { - p.mu.Lock() - defer p.mu.Unlock() - - var raw map[string]json.RawMessage - if err := json.Unmarshal(body, &raw); err != nil { - p.logger.Debug("contacts: malformed uids response", zap.Error(err)) - return nil, 0, 0, 0 - } - uidsRaw, ok := raw["uids"] - if !ok { - p.logger.Debug("contacts: uids response without uids key") - return nil, 0, 0, 0 - } - var uids []string - if err := json.Unmarshal(uidsRaw, &uids); err != nil { - p.logger.Debug("contacts: malformed uids list", zap.Error(err)) - return nil, 0, 0, 0 - } - if len(uids) > maxContactUIDs { - p.logger.Warn("contacts: uids response exceeds cap, refusing", - zap.Int("count", len(uids))) - return nil, 0, 0, 0 - } - - dir, err := p.deviceDir(deviceID) - if err != nil { - p.logger.Warn("contacts: bad device dir", zap.Error(err)) - return nil, 0, 0, 0 - } - idx := loadIndex(dir) - - seen := make(map[string]bool, len(uids)) - var toFetch []string - var added, updated int - for _, uid := range uids { - if uid == "" { - continue - } - seen[uid] = true - ts := coerceTimestamp(raw[uid]) - entry, known := idx[uid] - if !known { - added++ - toFetch = append(toFetch, uid) - } else if entry.Timestamp != ts { - updated++ - toFetch = append(toFetch, uid) - } - // Record the authoritative timestamp now (the vCards round - // carries no stamps); the summary fields arrive with the vCard. - entry.Timestamp = ts - idx[uid] = entry - } - - // Delete locally-known contacts the phone no longer reports — but - // never on an empty list: a buggy/empty response must not wipe - // the cache. - var deleted int - if len(uids) > 0 { - for uid := range idx { - if !seen[uid] { - if path, err := contactPath(dir, uid); err == nil { - _ = os.Remove(path) - } - delete(idx, uid) - deleted++ - } - } - } - if err := saveIndex(dir, idx); err != nil { - p.logger.Warn("contacts: failed to save index", zap.Error(err)) - } - return toFetch, added, updated, deleted - }() - - if len(toFetch) > 0 { - if err := p.requestVCards(dev, toFetch); err != nil { - p.logger.Warn("contacts: failed to request vcards", - zap.String("device_id", deviceID), - zap.Error(err)) - } - } - p.emit(deviceID, map[string]any{ - "phase": "uids", - "added": added, - "updated": updated, - "deleted": deleted, - "pending": len(toFetch), - }) -} - -// handleVCardsResponse stores one .vcf per UID, refreshes the index, and -// reports counts (never contact content) on the event bus. -func (p *ContactsPlugin) handleVCardsResponse(deviceID string, body []byte) { - p.mu.Lock() - defer p.mu.Unlock() - - var raw map[string]json.RawMessage - if err := json.Unmarshal(body, &raw); err != nil { - p.logger.Debug("contacts: malformed vcards response", zap.Error(err)) - return - } - uidsRaw, ok := raw["uids"] - if !ok { - p.logger.Debug("contacts: vcards response without uids key") - return - } - var uids []string - if err := json.Unmarshal(uidsRaw, &uids); err != nil { - p.logger.Debug("contacts: malformed vcards uids list", zap.Error(err)) - return - } - if len(uids) > maxContactUIDs { - p.logger.Warn("contacts: vcards response exceeds cap, refusing", - zap.Int("count", len(uids))) - return - } - - dir, err := p.deviceDir(deviceID) - if err != nil { - p.logger.Warn("contacts: bad device dir", zap.Error(err)) - return - } - idx := loadIndex(dir) - - var stored, skipped int - var total int64 - for _, uid := range uids { - if uid == "" { - continue - } - vRaw, ok := raw[uid] - if !ok { - continue - } - var vcard string - if err := json.Unmarshal(vRaw, &vcard); err != nil { - p.logger.Debug("contacts: non-string vcard, skipping", - zap.String("uid", sanitizeUID(uid))) - skipped++ - continue - } - if len(vcard) > maxVCardBytes { - p.logger.Warn("contacts: oversized vcard, skipping", - zap.String("uid", sanitizeUID(uid)), - zap.Int("bytes", len(vcard))) - skipped++ - continue - } - total += int64(len(vcard)) - if total > maxSyncBytes { - p.logger.Warn("contacts: sync exceeds total cap, stopping") - break - } - path, err := contactPath(dir, uid) - if err != nil { - skipped++ - continue - } - if err := os.WriteFile(path, []byte(vcard), 0600); err != nil { - p.logger.Warn("contacts: failed to store vcard", zap.Error(err)) - skipped++ - continue - } - name, phones, emails := parseVCard(vcard) - // The vCards round carries no per-UID stamps — keep the timestamp - // established by the uids round and mark the entry fetched. - entry := idx[uid] - entry.Fetched = true - entry.Name = name - entry.Phones = phones - entry.Emails = emails - idx[uid] = entry - stored++ - } - if err := saveIndex(dir, idx); err != nil { - p.logger.Warn("contacts: failed to save index", zap.Error(err)) - } - - p.emit(deviceID, map[string]any{ - "phase": "vcards", - "stored": stored, - "skipped": skipped, - }) -} - -// List returns cached contact summaries sorted by name. Entries whose -// vCard hasn't arrived yet are skipped; empty when never synced — absent -// means unknown, never a fabricated entry. -func (p *ContactsPlugin) List(deviceID string) []ContactSummary { - p.mu.Lock() - defer p.mu.Unlock() - - dir, err := p.cacheDir(deviceID) - if err != nil { - return nil - } - idx := loadIndex(dir) - out := make([]ContactSummary, 0, len(idx)) - for uid, entry := range idx { - if !entry.Fetched { - continue - } - out = append(out, ContactSummary{ - UID: uid, - Name: entry.Name, - Phones: entry.Phones, - Emails: entry.Emails, - Timestamp: entry.Timestamp, - }) - } - sort.Slice(out, func(i, j int) bool { - if out[i].Name != out[j].Name { - return out[i].Name < out[j].Name - } - return out[i].UID < out[j].UID - }) - return out -} - -// ForgetDevice deletes a device's cached contacts. Called on unpair: -// revoked trust drops the address book; re-pair re-syncs from scratch. -func (p *ContactsPlugin) ForgetDevice(deviceID string) error { - p.mu.Lock() - defer p.mu.Unlock() - - dir, err := p.cacheDir(deviceID) - if err != nil { - return err - } - if err := os.RemoveAll(dir); err != nil { - return fmt.Errorf("contacts: forget device: %w", err) - } - return nil -} - -func (p *ContactsPlugin) emit(deviceID string, payload map[string]any) { - if p.bus == nil { - return - } - p.bus.Publish(events.TypeContactsUpdated, deviceID, payload) -} - -// cleanDisplayValue strips control characters (terminal-escape injection -// via contact names is a classic) and truncates overlong fields. -func cleanDisplayValue(s string) string { - s = strings.Map(func(r rune) rune { - if r < 0x20 || r == 0x7f { - return -1 - } - return r - }, s) - s = strings.TrimSpace(s) - if len(s) > maxDisplayLen { - s = s[:maxDisplayLen] - } - return s -} - -// parseVCard extracts display fields from vCard 3.0 text with a stdlib -// line scan: unfold continuations, split Name:Value, drop parameters -// (TEL;TYPE=CELL:...). Returns the first FN and all TEL/EMAIL values. -func parseVCard(vcard string) (name string, phones, emails []string) { - // Normalize newlines, then unfold continuation lines (leading SP/HT). - raw := strings.ReplaceAll(vcard, "\r\n", "\n") - raw = strings.ReplaceAll(raw, "\r", "\n") - var lines []string - for _, line := range strings.Split(raw, "\n") { - if line == "" { - continue - } - if (line[0] == ' ' || line[0] == '\t') && len(lines) > 0 { - lines[len(lines)-1] += line[1:] - continue - } - lines = append(lines, line) - } - for _, line := range lines { - sep := strings.IndexByte(line, ':') - if sep < 0 { - continue - } - field := strings.ToUpper(line[:sep]) - if i := strings.IndexByte(field, ';'); i >= 0 { - field = field[:i] - } - value := cleanDisplayValue(line[sep+1:]) - if value == "" { - continue - } - switch field { - case "FN": - if name == "" { - name = value - } - case "TEL": - phones = append(phones, value) - case "EMAIL": - emails = append(emails, value) - } - } - return name, phones, emails -} diff --git a/internal/plugins/contacts/request.go b/internal/plugins/contacts/request.go new file mode 100644 index 0000000..c680e13 --- /dev/null +++ b/internal/plugins/contacts/request.go @@ -0,0 +1,78 @@ +package contacts + +import ( + "context" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" +) + +// Handle routes response packets; all parsing and disk I/O runs in a +// worker goroutine so Handle returns immediately (rule 9). +func (p *ContactsPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + switch pkt.Type { + case PacketTypeContactsResponseUIDs: + body := append([]byte(nil), pkt.Body...) + go p.handleUIDsResponse(dev, body) + return nil + case PacketTypeContactsResponseVCards: + body := append([]byte(nil), pkt.Body...) + devID := dev.ID() + go p.handleVCardsResponse(devID, body) + return nil + default: + return nil + } +} + +// RequestSync asks the phone for all contact UIDs and timestamps, which +// starts the sync round trips. Responses arrive async via Handle. +// +// The body must be an empty object, not null: stock implementations send +// `"body":{}` for bodyless requests, and at least one phone build aborts +// the whole link on an explicit null body (observed as an immediate RST +// after every connect that carried `"body":null`). +func (p *ContactsPlugin) RequestSync(dev device.Sender) error { + pkt, err := protocol.NewPacket(PacketTypeContactsRequestUIDs, map[string]any{}) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// requestVCards asks for vCards of the given UIDs, chunked to bound +// outbound packet size. +func (p *ContactsPlugin) requestVCards(dev device.Sender, uids []string) error { + for _, chunk := range chunkUIDs(uids, vcardChunkSize) { + body := map[string]any{"uids": chunk} + pkt, err := protocol.NewPacket(PacketTypeContactsRequestVCards, body) + if err != nil { + return err + } + if err := dev.Send(pkt); err != nil { + return err + } + } + return nil +} + +func chunkUIDs(uids []string, size int) [][]string { + var chunks [][]string + for len(uids) > 0 { + n := size + if len(uids) < n { + n = len(uids) + } + chunks = append(chunks, uids[:n]) + uids = uids[n:] + } + return chunks +} + +func (p *ContactsPlugin) emit(deviceID string, payload map[string]any) { + if p.bus == nil { + return + } + p.bus.Publish(events.TypeContactsUpdated, deviceID, payload) +} diff --git a/internal/plugins/contacts/store.go b/internal/plugins/contacts/store.go new file mode 100644 index 0000000..b4d0a9b --- /dev/null +++ b/internal/plugins/contacts/store.go @@ -0,0 +1,122 @@ +package contacts + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" +) + +// cacheDir resolves the cache dir for a device without creating it. +// The device ID comes from our own registry, but confine defensively. +func (p *ContactsPlugin) cacheDir(deviceID string) (string, error) { + safe := sanitizeUID(deviceID) + if safe == "" { + return "", fmt.Errorf("contacts: unusable device id") + } + dir := filepath.Join(p.baseDir, safe) + if rel, err := filepath.Rel(p.baseDir, dir); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + return "", fmt.Errorf("contacts: device dir escapes cache") + } + return dir, nil +} + +// deviceDir returns the cache dir for a device, creating it (0700). +func (p *ContactsPlugin) deviceDir(deviceID string) (string, error) { + dir, err := p.cacheDir(deviceID) + if err != nil { + return "", err + } + if err := os.MkdirAll(dir, 0700); err != nil { + return "", fmt.Errorf("contacts: create dir: %w", err) + } + return dir, nil +} + +// contactPath resolves the vcf file for a UID inside dir, confined to dir. +func contactPath(dir, uid string) (string, error) { + safe := sanitizeUID(uid) + if safe == "" { + return "", fmt.Errorf("contacts: unusable uid") + } + path := filepath.Join(dir, safe+".vcf") + if rel, err := filepath.Rel(dir, path); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + return "", fmt.Errorf("contacts: uid escapes cache") + } + return path, nil +} + +// loadIndex reads the sidecar index (missing file = empty cache, not error). +func loadIndex(dir string) map[string]indexEntry { + idx := make(map[string]indexEntry) + data, err := os.ReadFile(filepath.Join(dir, "index.json")) + if err != nil { + return idx + } + _ = json.Unmarshal(data, &idx) + if idx == nil { + idx = make(map[string]indexEntry) + } + return idx +} + +// saveIndex persists the sidecar index (0600). +func saveIndex(dir string, idx map[string]indexEntry) error { + data, err := json.MarshalIndent(idx, "", " ") + if err != nil { + return err + } + return os.WriteFile(filepath.Join(dir, "index.json"), data, 0600) +} + +// List returns cached contact summaries sorted by name. Entries whose +// vCard hasn't arrived yet are skipped; empty when never synced — absent +// means unknown, never a fabricated entry. +func (p *ContactsPlugin) List(deviceID string) []ContactSummary { + p.mu.Lock() + defer p.mu.Unlock() + + dir, err := p.cacheDir(deviceID) + if err != nil { + return nil + } + idx := loadIndex(dir) + out := make([]ContactSummary, 0, len(idx)) + for uid, entry := range idx { + if !entry.Fetched { + continue + } + out = append(out, ContactSummary{ + UID: uid, + Name: entry.Name, + Phones: entry.Phones, + Emails: entry.Emails, + Timestamp: entry.Timestamp, + }) + } + sort.Slice(out, func(i, j int) bool { + if out[i].Name != out[j].Name { + return out[i].Name < out[j].Name + } + return out[i].UID < out[j].UID + }) + return out +} + +// ForgetDevice deletes a device's cached contacts. Called on unpair: +// revoked trust drops the address book; re-pair re-syncs from scratch. +func (p *ContactsPlugin) ForgetDevice(deviceID string) error { + p.mu.Lock() + defer p.mu.Unlock() + + dir, err := p.cacheDir(deviceID) + if err != nil { + return err + } + if err := os.RemoveAll(dir); err != nil { + return fmt.Errorf("contacts: forget device: %w", err) + } + return nil +} diff --git a/internal/plugins/contacts/sync.go b/internal/plugins/contacts/sync.go new file mode 100644 index 0000000..5b8d433 --- /dev/null +++ b/internal/plugins/contacts/sync.go @@ -0,0 +1,231 @@ +package contacts + +import ( + "encoding/json" + "os" + "strconv" + "strings" + + "github.com/bethropolis/kcd/internal/device" + "go.uber.org/zap" +) + +// coerceTimestamp parses Android's string-encoded timestamps as well as +// plain JSON numbers (the header doc shows ints). Garbage yields 0, which +// forces a re-fetch — the safe direction, never a false "unchanged". +func coerceTimestamp(raw json.RawMessage) int64 { + var asAny any + if err := json.Unmarshal(raw, &asAny); err != nil { + return 0 + } + switch v := asAny.(type) { + case string: + n, err := strconv.ParseInt(strings.TrimSpace(v), 10, 64) + if err != nil { + return 0 + } + return n + case float64: + if v < 0 { + return 0 + } + return int64(v) + default: + return 0 + } +} + +// handleUIDsResponse diffs the phone's UID/timestamp list against the local +// cache, deletes stale entries, and requests vCards for new/changed ones. +// dev is the live sender (device.Send is channel-safe across goroutines). +func (p *ContactsPlugin) handleUIDsResponse(dev device.Sender, body []byte) { + deviceID := dev.ID() + + // Cache diff under lock; the vCard request round below does network + // I/O and must not hold it. + toFetch, added, updated, deleted := func() ([]string, int, int, int) { + p.mu.Lock() + defer p.mu.Unlock() + + var raw map[string]json.RawMessage + if err := json.Unmarshal(body, &raw); err != nil { + p.logger.Debug("contacts: malformed uids response", zap.Error(err)) + return nil, 0, 0, 0 + } + uidsRaw, ok := raw["uids"] + if !ok { + p.logger.Debug("contacts: uids response without uids key") + return nil, 0, 0, 0 + } + var uids []string + if err := json.Unmarshal(uidsRaw, &uids); err != nil { + p.logger.Debug("contacts: malformed uids list", zap.Error(err)) + return nil, 0, 0, 0 + } + if len(uids) > maxContactUIDs { + p.logger.Warn("contacts: uids response exceeds cap, refusing", + zap.Int("count", len(uids))) + return nil, 0, 0, 0 + } + + dir, err := p.deviceDir(deviceID) + if err != nil { + p.logger.Warn("contacts: bad device dir", zap.Error(err)) + return nil, 0, 0, 0 + } + idx := loadIndex(dir) + + seen := make(map[string]bool, len(uids)) + var toFetch []string + var added, updated int + for _, uid := range uids { + if uid == "" { + continue + } + seen[uid] = true + ts := coerceTimestamp(raw[uid]) + entry, known := idx[uid] + if !known { + added++ + toFetch = append(toFetch, uid) + } else if entry.Timestamp != ts { + updated++ + toFetch = append(toFetch, uid) + } + // Record the authoritative timestamp now (the vCards round + // carries no stamps); the summary fields arrive with the vCard. + entry.Timestamp = ts + idx[uid] = entry + } + + // Delete locally-known contacts the phone no longer reports — but + // never on an empty list: a buggy/empty response must not wipe + // the cache. + var deleted int + if len(uids) > 0 { + for uid := range idx { + if !seen[uid] { + if path, err := contactPath(dir, uid); err == nil { + _ = os.Remove(path) + } + delete(idx, uid) + deleted++ + } + } + } + if err := saveIndex(dir, idx); err != nil { + p.logger.Warn("contacts: failed to save index", zap.Error(err)) + } + return toFetch, added, updated, deleted + }() + + if len(toFetch) > 0 { + if err := p.requestVCards(dev, toFetch); err != nil { + p.logger.Warn("contacts: failed to request vcards", + zap.String("device_id", deviceID), + zap.Error(err)) + } + } + p.emit(deviceID, map[string]any{ + "phase": "uids", + "added": added, + "updated": updated, + "deleted": deleted, + "pending": len(toFetch), + }) +} + +// handleVCardsResponse stores one .vcf per UID, refreshes the index, and +// reports counts (never contact content) on the event bus. +func (p *ContactsPlugin) handleVCardsResponse(deviceID string, body []byte) { + p.mu.Lock() + defer p.mu.Unlock() + + var raw map[string]json.RawMessage + if err := json.Unmarshal(body, &raw); err != nil { + p.logger.Debug("contacts: malformed vcards response", zap.Error(err)) + return + } + uidsRaw, ok := raw["uids"] + if !ok { + p.logger.Debug("contacts: vcards response without uids key") + return + } + var uids []string + if err := json.Unmarshal(uidsRaw, &uids); err != nil { + p.logger.Debug("contacts: malformed vcards uids list", zap.Error(err)) + return + } + if len(uids) > maxContactUIDs { + p.logger.Warn("contacts: vcards response exceeds cap, refusing", + zap.Int("count", len(uids))) + return + } + + dir, err := p.deviceDir(deviceID) + if err != nil { + p.logger.Warn("contacts: bad device dir", zap.Error(err)) + return + } + idx := loadIndex(dir) + + var stored, skipped int + var total int64 + for _, uid := range uids { + if uid == "" { + continue + } + vRaw, ok := raw[uid] + if !ok { + continue + } + var vcard string + if err := json.Unmarshal(vRaw, &vcard); err != nil { + p.logger.Debug("contacts: non-string vcard, skipping", + zap.String("uid", sanitizeUID(uid))) + skipped++ + continue + } + if len(vcard) > maxVCardBytes { + p.logger.Warn("contacts: oversized vcard, skipping", + zap.String("uid", sanitizeUID(uid)), + zap.Int("bytes", len(vcard))) + skipped++ + continue + } + total += int64(len(vcard)) + if total > maxSyncBytes { + p.logger.Warn("contacts: sync exceeds total cap, stopping") + break + } + path, err := contactPath(dir, uid) + if err != nil { + skipped++ + continue + } + if err := os.WriteFile(path, []byte(vcard), 0600); err != nil { + p.logger.Warn("contacts: failed to store vcard", zap.Error(err)) + skipped++ + continue + } + name, phones, emails := parseVCard(vcard) + // The vCards round carries no per-UID stamps — keep the timestamp + // established by the uids round and mark the entry fetched. + entry := idx[uid] + entry.Fetched = true + entry.Name = name + entry.Phones = phones + entry.Emails = emails + idx[uid] = entry + stored++ + } + if err := saveIndex(dir, idx); err != nil { + p.logger.Warn("contacts: failed to save index", zap.Error(err)) + } + + p.emit(deviceID, map[string]any{ + "phase": "vcards", + "stored": stored, + "skipped": skipped, + }) +} diff --git a/internal/plugins/contacts/types.go b/internal/plugins/contacts/types.go new file mode 100644 index 0000000..80e889b --- /dev/null +++ b/internal/plugins/contacts/types.go @@ -0,0 +1,114 @@ +package contacts + +import ( + "os" + "path/filepath" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "go.uber.org/zap" +) + +// KDE Connect contacts packet types (see kdeconnect-kde +// plugins/contacts/contactsplugin.h and the Android ContactsPlugin). +const ( + PacketTypeContactsRequestUIDs = "kdeconnect.contacts.request_all_uids_timestamps" + PacketTypeContactsRequestVCards = "kdeconnect.contacts.request_vcards_by_uid" + PacketTypeContactsResponseUIDs = "kdeconnect.contacts.response_uids_timestamps" + PacketTypeContactsResponseVCards = "kdeconnect.contacts.response_vcards" +) + +const ( + // maxContactUIDs bounds a single uids response (address books are in + // the hundreds; this is two orders of magnitude of headroom). + maxContactUIDs = 20000 + // maxVCardBytes caps one vCard (text contact cards are ~1KB). + maxVCardBytes = 64 * 1024 + // maxSyncBytes caps a whole vcards response before anything hits disk. + maxSyncBytes = 64 * 1024 * 1024 + // vcardChunkSize bounds our outbound request packets. + vcardChunkSize = 100 + // maxDisplayLen truncates contact fields shown to clients. + maxDisplayLen = 256 +) + +// ContactSummary is the parsed, display-safe view of one contact. +type ContactSummary struct { + UID string `json:"uid"` + Name string `json:"name"` + Phones []string `json:"phones,omitempty"` + Emails []string `json:"emails,omitempty"` + Timestamp int64 `json:"timestamp"` +} + +// indexEntry is the persisted per-contact record (index.json sidecar, so +// listing never reparses hundreds of vCards). Fetched marks that the vCard +// round completed; entries with only a timestamp are pending fetch. +type indexEntry struct { + Timestamp int64 `json:"timestamp"` + Fetched bool `json:"fetched"` + Name string `json:"name"` + Phones []string `json:"phones,omitempty"` + Emails []string `json:"emails,omitempty"` +} + +// ContactsPlugin syncs the phone address book: it requests UID/timestamp +// lists, fetches vCards for new or changed contacts, caches them per +// device, and deletes stale entries. +type ContactsPlugin struct { + bus *events.Bus + logger *zap.Logger + baseDir string + + // mu serializes sync processing across devices. Syncs are rare; + // one lock avoids per-device lock bookkeeping. + mu sync.Mutex +} + +// NewContactsPlugin creates a contacts plugin caching under +// $XDG_DATA_HOME/kcd/contacts (0600 files, 0700 dirs). +func NewContactsPlugin(bus *events.Bus, logger *zap.Logger, cacheDirs ...string) *ContactsPlugin { + dataHome := os.Getenv("XDG_DATA_HOME") + if dataHome == "" { + home, _ := os.UserHomeDir() + dataHome = filepath.Join(home, ".local", "share") + } + baseDir := filepath.Join(dataHome, "kcd", "contacts") + if len(cacheDirs) > 0 && cacheDirs[0] != "" { + baseDir = cacheDirs[0] + } + _ = os.MkdirAll(baseDir, 0700) + + return &ContactsPlugin{ + bus: bus, + logger: logger.With(zap.String("plugin", "contacts")), + baseDir: baseDir, + } +} + +func (p *ContactsPlugin) Name() string { return "Contacts" } + +func (p *ContactsPlugin) Timeout() time.Duration { return 5 * time.Second } + +func (p *ContactsPlugin) IncomingTypes() []string { + return []string{PacketTypeContactsResponseUIDs, PacketTypeContactsResponseVCards} +} + +func (p *ContactsPlugin) OutgoingTypes() []string { + return []string{PacketTypeContactsRequestUIDs, PacketTypeContactsRequestVCards} +} + +// OnConnect starts a sync for freshly connected paired devices, mirroring +// upstream's connected() -> synchronizeRemoteWithLocal(). +func (p *ContactsPlugin) OnConnect(dev device.Sender) { + if dev.State() != device.StatePaired { + return + } + if err := p.RequestSync(dev); err != nil { + p.logger.Debug("contacts: initial sync request failed", zap.Error(err)) + } +} + +func (p *ContactsPlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/contacts/vcard.go b/internal/plugins/contacts/vcard.go new file mode 100644 index 0000000..927bd67 --- /dev/null +++ b/internal/plugins/contacts/vcard.go @@ -0,0 +1,77 @@ +package contacts + +import ( + "regexp" + "strings" +) + +// uidSafeChars restricts phone-provided contact UIDs (Android LOOKUP_KEYs) +// to filename-safe characters so they can't escape the cache directory. +var uidSafeChars = regexp.MustCompile(`[^a-zA-Z0-9._-]`) + +// sanitizeUID strips everything but alphanumerics, dot, underscore and +// hyphen. Empty results are rejected by the caller. +func sanitizeUID(uid string) string { + return uidSafeChars.ReplaceAllString(uid, "_") +} + +// cleanDisplayValue strips control characters (terminal-escape injection +// via contact names is a classic) and truncates overlong fields. +func cleanDisplayValue(s string) string { + s = strings.Map(func(r rune) rune { + if r < 0x20 || r == 0x7f { + return -1 + } + return r + }, s) + s = strings.TrimSpace(s) + if len(s) > maxDisplayLen { + s = s[:maxDisplayLen] + } + return s +} + +// parseVCard extracts display fields from vCard 3.0 text with a stdlib +// line scan: unfold continuations, split Name:Value, drop parameters +// (TEL;TYPE=CELL:...). Returns the first FN and all TEL/EMAIL values. +func parseVCard(vcard string) (name string, phones, emails []string) { + // Normalize newlines, then unfold continuation lines (leading SP/HT). + raw := strings.ReplaceAll(vcard, "\r\n", "\n") + raw = strings.ReplaceAll(raw, "\r", "\n") + var lines []string + for _, line := range strings.Split(raw, "\n") { + if line == "" { + continue + } + if (line[0] == ' ' || line[0] == '\t') && len(lines) > 0 { + lines[len(lines)-1] += line[1:] + continue + } + lines = append(lines, line) + } + for _, line := range lines { + sep := strings.IndexByte(line, ':') + if sep < 0 { + continue + } + field := strings.ToUpper(line[:sep]) + if i := strings.IndexByte(field, ';'); i >= 0 { + field = field[:i] + } + value := cleanDisplayValue(line[sep+1:]) + if value == "" { + continue + } + switch field { + case "FN": + if name == "" { + name = value + } + case "TEL": + phones = append(phones, value) + case "EMAIL": + emails = append(emails, value) + } + } + return name, phones, emails +} diff --git a/internal/plugins/mousepad/dispatch.go b/internal/plugins/mousepad/dispatch.go new file mode 100644 index 0000000..48ef425 --- /dev/null +++ b/internal/plugins/mousepad/dispatch.go @@ -0,0 +1,63 @@ +package mousepad + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" +) + +func (p *MousepadPlugin) Handle(_ context.Context, _ device.Sender, pkt *protocol.Packet) error { + var body MousepadBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return fmt.Errorf("mousepad: decode body: %w", err) + } + + isPointerMove := (body.Dx != 0 || body.Dy != 0) && !body.SingleClick && !body.DoubleClick && + !body.RightClick && !body.MiddleClick && !body.SingleHold && !body.SingleRel && + body.Key == "" && body.SpecialKey == 0 + + if isPointerMove { + // Drop stale frame if worker hasn't consumed the last one yet. + select { + case p.moveCh <- body: + default: + select { + case <-p.moveCh: // drain + default: + } + p.moveCh <- body + } + } else { + p.eventCh <- body + } + return nil +} + +// worker is the single persistent goroutine that processes all mousepad events. +func (p *MousepadPlugin) worker() { + for { + select { + case <-p.ctx.Done(): + return + case body, ok := <-p.moveCh: + if !ok { + return + } + p.handleMove(body) + case body, ok := <-p.eventCh: + if !ok { + return + } + p.handleEvent(body) + } + } +} + +// Close stops the background worker goroutine. +func (p *MousepadPlugin) Close() error { + p.cancel() + return nil +} diff --git a/internal/plugins/mousepad/mousepad.go b/internal/plugins/mousepad/input.go similarity index 53% rename from internal/plugins/mousepad/mousepad.go rename to internal/plugins/mousepad/input.go index a7b7388..abd7362 100644 --- a/internal/plugins/mousepad/mousepad.go +++ b/internal/plugins/mousepad/input.go @@ -2,174 +2,15 @@ package mousepad import ( "context" - "encoding/json" - "fmt" - "os" "os/exec" "strconv" - "time" "github.com/bendahl/uinput" - "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/protocol" "go.uber.org/zap" ) -type MousepadPlugin struct { - logger *zap.Logger - cfg config.MousepadConfig - useYdotool bool - useUinput bool - isWayland bool - - // uinput devices - mouse uinput.Mouse - keyboard uinput.Keyboard - - moveCh chan MousepadBody // capacity 1, older frames dropped - eventCh chan MousepadBody // capacity 64, clicks + keys - - ctx context.Context - cancel context.CancelFunc -} - -func NewMousepadPlugin(cfg config.MousepadConfig, logger *zap.Logger) *MousepadPlugin { - ctx, cancel := context.WithCancel(context.Background()) - isWayland := os.Getenv("WAYLAND_DISPLAY") != "" - p := &MousepadPlugin{ - logger: logger.With(zap.String("plugin", "mousepad")), - cfg: cfg, - isWayland: isWayland, - moveCh: make(chan MousepadBody, 1), - eventCh: make(chan MousepadBody, 64), - ctx: ctx, - cancel: cancel, - } - - // Try uinput first if auto or explicit - if cfg.Backend == "auto" || cfg.Backend == "uinput" { - if err := p.initUinput(); err != nil { - p.logger.Warn("uinput initialization failed, falling back to legacy backends", zap.Error(err)) - } else { - p.useUinput = true - p.logger.Info("uinput initialized successfully") - } - } - - if !p.useUinput { - switch cfg.Backend { - case "ydotool": - p.useYdotool = true - case "xdotool": - p.useYdotool = false - default: // auto - p.useYdotool = os.Getenv("WAYLAND_DISPLAY") != "" - } - } - - go p.worker() - return p -} - -func (p *MousepadPlugin) initUinput() error { - m, err := uinput.CreateMouse("/dev/uinput", []byte("kcd-mouse")) - if err != nil { - return fmt.Errorf("create mouse: %w", err) - } - p.mouse = m - - k, err := uinput.CreateKeyboard("/dev/uinput", []byte("kcd-keyboard")) - if err != nil { - m.Close() - return fmt.Errorf("create keyboard: %w", err) - } - p.keyboard = k - - return nil -} - -// MousepadBody represents the exact spec Android sends. -type MousepadBody struct { - Dx float64 `json:"dx"` - Dy float64 `json:"dy"` - X float64 `json:"x"` - Y float64 `json:"y"` - SingleClick bool `json:"singleclick"` - DoubleClick bool `json:"doubleclick"` - MiddleClick bool `json:"middleclick"` - RightClick bool `json:"rightclick"` - SingleHold bool `json:"singlehold"` - SingleRel bool `json:"singlerelease"` - Scroll bool `json:"scroll"` - Key string `json:"key"` - SpecialKey int `json:"specialKey"` - Shift bool `json:"shift"` - Ctrl bool `json:"ctrl"` - Alt bool `json:"alt"` - Super bool `json:"super"` -} - -func (p *MousepadPlugin) Name() string { return "Mousepad" } -func (p *MousepadPlugin) Timeout() time.Duration { return 2 * time.Second } -func (p *MousepadPlugin) IncomingTypes() []string { return []string{"kdeconnect.mousepad.request"} } -func (p *MousepadPlugin) OutgoingTypes() []string { - return []string{"kdeconnect.mousepad.keyboardstate"} -} - -func (p *MousepadPlugin) Handle(_ context.Context, _ device.Sender, pkt *protocol.Packet) error { - var body MousepadBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return fmt.Errorf("mousepad: decode body: %w", err) - } - - isPointerMove := (body.Dx != 0 || body.Dy != 0) && !body.SingleClick && !body.DoubleClick && - !body.RightClick && !body.MiddleClick && !body.SingleHold && !body.SingleRel && - body.Key == "" && body.SpecialKey == 0 - - if isPointerMove { - // Drop stale frame if worker hasn't consumed the last one yet. - select { - case p.moveCh <- body: - default: - select { - case <-p.moveCh: // drain - default: - } - p.moveCh <- body - } - } else { - p.eventCh <- body - } - return nil -} - -// worker is the single persistent goroutine that processes all mousepad events. -func (p *MousepadPlugin) worker() { - for { - select { - case <-p.ctx.Done(): - return - case body, ok := <-p.moveCh: - if !ok { - return - } - p.handleMove(body) - case body, ok := <-p.eventCh: - if !ok { - return - } - p.handleEvent(body) - } - } -} - -// Close stops the background worker goroutine. -func (p *MousepadPlugin) Close() error { - p.cancel() - return nil -} - func (p *MousepadPlugin) handleMove(body MousepadBody) { if body.Dx == 0 && body.Dy == 0 { return diff --git a/internal/plugins/mousepad/types.go b/internal/plugins/mousepad/types.go new file mode 100644 index 0000000..ea22b33 --- /dev/null +++ b/internal/plugins/mousepad/types.go @@ -0,0 +1,113 @@ +package mousepad + +import ( + "context" + "fmt" + "os" + "time" + + "github.com/bendahl/uinput" + "github.com/bethropolis/kcd/internal/config" + "go.uber.org/zap" +) + +type MousepadPlugin struct { + logger *zap.Logger + cfg config.MousepadConfig + useYdotool bool + useUinput bool + isWayland bool + + // uinput devices + mouse uinput.Mouse + keyboard uinput.Keyboard + + moveCh chan MousepadBody // capacity 1, older frames dropped + eventCh chan MousepadBody // capacity 64, clicks + keys + + ctx context.Context + cancel context.CancelFunc +} + +func NewMousepadPlugin(cfg config.MousepadConfig, logger *zap.Logger) *MousepadPlugin { + ctx, cancel := context.WithCancel(context.Background()) + isWayland := os.Getenv("WAYLAND_DISPLAY") != "" + p := &MousepadPlugin{ + logger: logger.With(zap.String("plugin", "mousepad")), + cfg: cfg, + isWayland: isWayland, + moveCh: make(chan MousepadBody, 1), + eventCh: make(chan MousepadBody, 64), + ctx: ctx, + cancel: cancel, + } + + // Try uinput first if auto or explicit + if cfg.Backend == "auto" || cfg.Backend == "uinput" { + if err := p.initUinput(); err != nil { + p.logger.Warn("uinput initialization failed, falling back to legacy backends", zap.Error(err)) + } else { + p.useUinput = true + p.logger.Info("uinput initialized successfully") + } + } + + if !p.useUinput { + switch cfg.Backend { + case "ydotool": + p.useYdotool = true + case "xdotool": + p.useYdotool = false + default: // auto + p.useYdotool = os.Getenv("WAYLAND_DISPLAY") != "" + } + } + + go p.worker() + return p +} + +func (p *MousepadPlugin) initUinput() error { + m, err := uinput.CreateMouse("/dev/uinput", []byte("kcd-mouse")) + if err != nil { + return fmt.Errorf("create mouse: %w", err) + } + p.mouse = m + + k, err := uinput.CreateKeyboard("/dev/uinput", []byte("kcd-keyboard")) + if err != nil { + m.Close() + return fmt.Errorf("create keyboard: %w", err) + } + p.keyboard = k + + return nil +} + +// MousepadBody represents the exact spec Android sends. +type MousepadBody struct { + Dx float64 `json:"dx"` + Dy float64 `json:"dy"` + X float64 `json:"x"` + Y float64 `json:"y"` + SingleClick bool `json:"singleclick"` + DoubleClick bool `json:"doubleclick"` + MiddleClick bool `json:"middleclick"` + RightClick bool `json:"rightclick"` + SingleHold bool `json:"singlehold"` + SingleRel bool `json:"singlerelease"` + Scroll bool `json:"scroll"` + Key string `json:"key"` + SpecialKey int `json:"specialKey"` + Shift bool `json:"shift"` + Ctrl bool `json:"ctrl"` + Alt bool `json:"alt"` + Super bool `json:"super"` +} + +func (p *MousepadPlugin) Name() string { return "Mousepad" } +func (p *MousepadPlugin) Timeout() time.Duration { return 2 * time.Second } +func (p *MousepadPlugin) IncomingTypes() []string { return []string{"kdeconnect.mousepad.request"} } +func (p *MousepadPlugin) OutgoingTypes() []string { + return []string{"kdeconnect.mousepad.keyboardstate"} +} diff --git a/internal/plugins/mpris/art.go b/internal/plugins/mpris/art.go new file mode 100644 index 0000000..2ca537f --- /dev/null +++ b/internal/plugins/mpris/art.go @@ -0,0 +1,165 @@ +package mpris + +import ( + "context" + "net/url" + "os" + "strings" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/plugins/share" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *MPRISPlugin) sendAlbumArt(ctx context.Context, dev device.Sender, player, artUrl string) { + p.mu.Lock() + reqKey := dev.ID() + "|" + artUrl + if lastReq, exists := p.artRequests[reqKey]; exists && time.Since(lastReq) < 5*time.Second { + p.mu.Unlock() + return + } + p.artRequests[reqKey] = time.Now() + p.mu.Unlock() + + cleanUrl := artUrl + if idx := strings.LastIndex(cleanUrl, "?t="); idx != -1 { + cleanUrl = cleanUrl[:idx] + } + + filePath := strings.TrimPrefix(cleanUrl, "file://") + if unescaped, err := url.PathUnescape(filePath); err == nil { + filePath = unescaped + } + + f, err := os.Open(filePath) + if err != nil { + p.logger.Debug("mpris: album art file not found", zap.String("path", filePath), zap.Error(err)) + return + } + stat, err := f.Stat() + f.Close() + if err != nil { + return + } + + var shareCfg config.ShareConfig + shareCfg.Defaults() + + ln, port, err := share.ListenSideChannel(ctx, shareCfg, p.tlsConfig) + if err != nil { + return + } + + go func() { + _ = share.AcceptAndSend(ln, filePath, p.tlsConfig, dev.ID(), cert.PinnedFingerprint(dev.PeerCert()), 10*time.Second, nil, p.logger) + }() + + pkt, err := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ + "transferringAlbumArt": true, + "player": player, + "albumArtUrl": artUrl, + }) + if err == nil { + pkt.PayloadSize = stat.Size() + pkt.PayloadTransferInfo = &protocol.TransferInfo{Port: port} + dev.Send(pkt) + } +} + +// requestAlbumArt asks the remote device to stream the album art bytes +// referenced by a kdeconnect:/artUri URI over a side channel. +func (p *MPRISPlugin) requestAlbumArt(dev device.Sender, player, artUrl string) { + if player == "" || artUrl == "" { + return + } + p.mu.Lock() + reqKey := "req|" + dev.ID() + "|" + artUrl + if lastReq, exists := p.artRequests[reqKey]; exists && time.Since(lastReq) < 10*time.Second { + p.mu.Unlock() + return + } + p.artRequests[reqKey] = time.Now() + p.mu.Unlock() + + body := MPRISRequest{ + Player: player, + AlbumArtUrl: artUrl, + } + pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) + if err != nil { + return + } + if err := dev.Send(pkt); err != nil { + p.logger.Debug("mpris: album art request failed", zap.Error(err)) + } +} + +// receiveAlbumArt streams an inbound album art payload into the cache +// and re-publishes the device state with a resolved file:// URL so watch +// clients and kcd mpris status surface the loadable location. +func (p *MPRISPlugin) receiveAlbumArt(_ context.Context, dev device.Sender, player, artUrl string, size int64, port int) { + remoteIP := dev.RemoteIP() + if remoteIP == nil { + return + } + if p.artCache == nil || size <= 0 || size > maxAlbumArtBytes { + return + } + + p.logger.Debug("mpris: receiving album art from remote", + zap.String("player", player), + zap.String("device_id", dev.ID()), + zap.Int64("size", size)) + + tmp, err := os.CreateTemp(p.artCache.Dir(), ".art-*") + if err != nil { + p.logger.Warn("mpris: failed to create temp file for album art", zap.Error(err)) + return + } + tmpPath := tmp.Name() + tmp.Close() + defer os.Remove(tmpPath) + + // The Handle ctx is canceled as soon as Handle returns; use an + // independent context so the side-channel dial isn't aborted. + dlCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := share.ReceiveSideChannel(dlCtx, remoteIP, port, size, tmpPath, p.tlsConfig, cert.PinnedFingerprint(dev.PeerCert()), nil, p.logger); err != nil { + p.logger.Warn("mpris: album art transfer failed", zap.Error(err)) + return + } + + fileURL, err := p.artCache.Commit(artUrl, tmpPath) + if err != nil { + p.logger.Warn("mpris: failed to cache album art", zap.Error(err)) + return + } + + p.stampAlbumArt(dev.ID(), player, artUrl, fileURL) +} + +// stampAlbumArt publishes a resolved file:// URL for a completed album-art +// fetch — but only if the device's current track still wants exactly this +// art. The side-channel download can take up to 30s, during which the track +// (or the active player) may have changed; stamping blindly would show the +// old track's cover on the new track. On mismatch the bytes stay in the art +// cache and the new track's own art request fulfills it. +func (p *MPRISPlugin) stampAlbumArt(deviceID, player, artUrl, fileURL string) { + p.mu.Lock() + state := p.remoteStates[deviceID] + if state == nil || state.AlbumArtUrl != artUrl || state.Player != player { + p.mu.Unlock() + return + } + state = state.DeepCopy() + state.AlbumArtUrl = fileURL + p.mu.Unlock() + if p.bus != nil { + p.bus.Publish(events.TypeMprisUpdate, deviceID, state) + } +} diff --git a/internal/plugins/mpris/handle.go b/internal/plugins/mpris/handle.go new file mode 100644 index 0000000..0ceaa75 --- /dev/null +++ b/internal/plugins/mpris/handle.go @@ -0,0 +1,176 @@ +package mpris + +import ( + "context" + "encoding/json" + "strings" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *MPRISPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + p.mu.Lock() + if _, exists := p.devices[dev.ID()]; !exists { + p.devices[dev.ID()] = dev + } + p.mu.Unlock() + + var body MPRISRequest + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return err + } + + if body.AlbumArtUrl != "" && strings.HasPrefix(body.AlbumArtUrl, "file://") { + go p.sendAlbumArt(ctx, dev, body.Player, body.AlbumArtUrl) + return nil + } + + // Inbound album art payload — the phone responds to requestAlbumArt + // with a side-channel transfer carrying the art bytes. + if body.TransferringAlbumArt && pkt.PayloadSize > 0 && pkt.PayloadTransferInfo != nil { + go p.receiveAlbumArt(ctx, dev, body.Player, body.AlbumArtUrl, + pkt.PayloadSize, pkt.PayloadTransferInfo.Port) + return nil + } + + if body.RequestPlayerList { + return p.sendPlayerList(dev) + } + + // Incoming playerList from remote device — prune players that no longer + // exist (their media session was destroyed) and request fresh status for + // the ones still around. + if body.PlayerList != nil { + p.logger.Debug("mpris: received player list from remote", zap.Strings("players", body.PlayerList)) + pruned := false + p.mu.Lock() + if prev := p.remoteStates[dev.ID()]; prev != nil { + inList := false + for _, name := range body.PlayerList { + if name == prev.Player { + inList = true + break + } + } + if !inList { + delete(p.remoteStates, dev.ID()) + delete(p.remoteStateTimes, dev.ID()) + delete(p.positionTrackers, dev.ID()) + pruned = true + } + } + p.mu.Unlock() + + if pruned && p.bus != nil { + // The tracked player's session is gone — emit an empty update so + // watchers fall back to "no media playing" for this device. + p.bus.Publish(events.TypeMprisUpdate, dev.ID(), &NowPlaying{}) + } + + for _, player := range body.PlayerList { + player := player + go p.requestPlayerStatus(dev, player) + } + return nil + } + + if body.Player == "" { + return nil + } + + if body.RequestNowPlaying || body.RequestVolume { + go func() { + if state, err := p.playerState(body.Player); err == nil { + p.broadcast(state) + } + }() + return nil + } + + if body.Action != "" || body.Seek != nil || body.SetPosition != nil || body.SetVolume != nil || body.SetShuffle != nil || body.SetLoopStatus != "" { + go p.handleAction(body.Player, body.Action, body.Seek, body.SetPosition, body.SetVolume, body.SetShuffle, body.SetLoopStatus) + return nil + } + + // Incoming NowPlaying state update from a remote device + if body.Title != "" || body.Artist != "" || body.Album != "" || body.IsPlaying || body.PlaybackStatus != "" { + state := &NowPlaying{ + Player: body.Player, + Title: body.Title, + Artist: body.Artist, + Album: body.Album, + AlbumArtUrl: body.AlbumArtUrl, + Url: body.Url, + Length: body.Length, + Pos: body.Pos, + IsPlaying: body.IsPlaying, + Volume: body.Volume, + CanControl: body.CanControl, + CanGoNext: body.CanGoNext, + CanGoPrevious: body.CanGoPrevious, + CanPause: body.CanPause, + CanPlay: body.CanPlay, + CanSeek: body.CanSeek, + PlaybackStatus: body.PlaybackStatus, + Shuffle: body.Shuffle, + LoopStatus: body.LoopStatus, + } + p.mu.Lock() + tracker, ok := p.positionTrackers[dev.ID()] + if !ok { + tracker = &remotePositionTracker{} + p.positionTrackers[dev.ID()] = tracker + } + tracker.lastPosition = body.Pos + tracker.lastPositionAt = time.Now() + tracker.playing = body.IsPlaying + state.PosAnchorMs = tracker.lastPositionAt.UnixMilli() + shouldPublish := shouldPublishRemoteState(p.remoteStates[dev.ID()], state) + p.remoteStates[dev.ID()] = state + p.remoteStateTimes[dev.ID()] = time.Now() + p.mu.Unlock() + + // Request art bytes from the phone when the advertsed album art is + // a kdeconnect:// URI we have not cached yet. Resolve any already + // cached art before publishing so watch clients get a loadable URL. + if p.artCache != nil && p.artCache.Resolve(state.AlbumArtUrl) == "" { + go p.requestAlbumArt(dev, state.Player, state.AlbumArtUrl) + } + if shouldPublish && p.bus != nil { + pub := state.DeepCopy() + if p.artCache != nil { + if resolved := p.artCache.Resolve(pub.AlbumArtUrl); resolved != "" { + pub.AlbumArtUrl = resolved + } else if pub.AlbumArtUrl != "" && strings.HasPrefix(pub.AlbumArtUrl, "kdeconnect:") { + // Art still downloading: publish an empty URL with the + // pending flag instead of an unloadable kdeconnect:/ + // URI. The arrival re-publish carries file://. + pub.AlbumArtUrl = "" + pub.ArtPending = true + } + } + p.bus.Publish(events.TypeMprisUpdate, dev.ID(), pub) + } + return nil + } + + return nil +} + +func shouldPublishRemoteState(last, current *NowPlaying) bool { + if last == nil { + return true + } + return last.Player != current.Player || + last.Title != current.Title || + last.Artist != current.Artist || + last.Album != current.Album || + last.AlbumArtUrl != current.AlbumArtUrl || + last.PlaybackStatus != current.PlaybackStatus || + last.IsPlaying != current.IsPlaying || + last.Volume != current.Volume +} diff --git a/internal/plugins/mpris/local.go b/internal/plugins/mpris/local.go new file mode 100644 index 0000000..68aaa64 --- /dev/null +++ b/internal/plugins/mpris/local.go @@ -0,0 +1,179 @@ +package mpris + +import ( + "strings" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *MPRISPlugin) sendPlayerList(dev device.Sender) error { + p.mu.RLock() + displayNames := make([]string, 0, len(p.players)) + for name := range p.players { + displayNames = append(displayNames, name) + } + p.mu.RUnlock() + + if displayNames == nil { + displayNames = []string{} + } + + p.logger.Debug("mpris: sending player list", zap.Strings("players", displayNames)) + + pkt, err := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ + "playerList": displayNames, + "supportAlbumArtPayload": true, + }) + if err != nil { + return err + } + + go func() { + for _, name := range displayNames { + if state, err := p.playerState(name); err == nil { + p.broadcast(state) + } + } + }() + + return dev.Send(pkt) +} + +func (p *MPRISPlugin) sendPlayerListBroadcast() { + p.mu.RLock() + displayNames := make([]string, 0, len(p.players)) + for name := range p.players { + displayNames = append(displayNames, name) + } + p.mu.RUnlock() + + if displayNames == nil { + displayNames = []string{} + } + + pkt, _ := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ + "playerList": displayNames, + "supportAlbumArtPayload": true, + }) + + p.mu.RLock() + defer p.mu.RUnlock() + for _, dev := range p.devices { + if dev.IsConnected() { + _ = dev.Send(pkt) + } + } +} + +func (p *MPRISPlugin) broadcast(state *NowPlaying) { + pkt, err := protocol.NewPacket("kdeconnect.mpris", state) + if err != nil { + return + } + + p.mu.RLock() + defer p.mu.RUnlock() + for _, dev := range p.devices { + if dev.IsConnected() { + _ = dev.Send(pkt) + } + } +} + +func (p *MPRISPlugin) addPlayer(busName, uniqueName, displayName, shortName string) { + p.mu.Lock() + p.players[displayName] = &trackedPlayer{ + busName: busName, + uniqueName: uniqueName, + displayName: displayName, + shortName: shortName, + } + p.mu.Unlock() + + p.logger.Debug("mpris: added player", zap.String("displayName", displayName), zap.String("busName", busName)) + + if state, err := p.playerState(displayName); err == nil { + p.mu.Lock() + p.lastStates[displayName] = state + p.mu.Unlock() + p.broadcast(state) + } + + p.sendPlayerListBroadcast() +} + +func (p *MPRISPlugin) removePlayer(displayName string) { + p.mu.Lock() + delete(p.players, displayName) + delete(p.lastTracks, displayName) + delete(p.lastStates, displayName) + p.mu.Unlock() + + p.logger.Debug("mpris: removed player", zap.String("displayName", displayName)) + + p.sendPlayerListBroadcast() +} + +func (p *MPRISPlugin) resolvePlayer(displayName string) *trackedPlayer { + if displayName == "" { + return nil + } + p.mu.RLock() + defer p.mu.RUnlock() + for _, pl := range p.players { + if pl.displayName == displayName || pl.shortName == displayName || strings.EqualFold(pl.shortName, displayName) { + return pl + } + } + return nil +} + +func (p *MPRISPlugin) DebugStatus() *DebugStatus { + p.mu.RLock() + watching := p.watching + devCount := len(p.devices) + playerMappings := make(map[string]string, len(p.players)) + playerList := make([]*trackedPlayer, 0, len(p.players)) + for _, pl := range p.players { + playerMappings[pl.displayName] = pl.busName + playerList = append(playerList, pl) + } + p.mu.RUnlock() + + var players []DebugPlayerInfo + for _, pl := range playerList { + info := DebugPlayerInfo{ + DisplayName: pl.displayName, + BusName: pl.busName, + ShortName: pl.shortName, + } + if state, err := p.playerState(pl.displayName); err == nil { + info.Title = state.Title + info.Artist = state.Artist + info.Album = state.Album + info.PlaybackStatus = state.PlaybackStatus + info.IsPlaying = state.IsPlaying + info.Volume = state.Volume + info.Pos = state.Pos + info.Length = state.Length + info.AlbumArtUrl = state.AlbumArtUrl + info.CanSeek = state.CanSeek + info.CanGoNext = state.CanGoNext + info.CanGoPrevious = state.CanGoPrevious + info.CanPlay = state.CanPlay + info.CanPause = state.CanPause + } else { + info.Error = err.Error() + } + players = append(players, info) + } + + return &DebugStatus{ + WatcherRunning: watching, + DeviceCount: devCount, + Players: players, + PlayerMappings: playerMappings, + } +} diff --git a/internal/plugins/mpris/mpris.go b/internal/plugins/mpris/mpris.go index c9d5b0e..88d3f4a 100644 --- a/internal/plugins/mpris/mpris.go +++ b/internal/plugins/mpris/mpris.go @@ -3,19 +3,11 @@ package mpris import ( "context" "crypto/tls" - "encoding/json" - "net/url" - "os" - "strings" "sync" "time" - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/config" "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/plugins/share" - "github.com/bethropolis/kcd/internal/protocol" "github.com/godbus/dbus/v5" "go.uber.org/zap" ) @@ -115,915 +107,3 @@ func (p *MPRISPlugin) IncomingTypes() []string { func (p *MPRISPlugin) OutgoingTypes() []string { return []string{"kdeconnect.mpris", "kdeconnect.mpris.request"} } - -type MPRISRequest struct { - // Request fields - RequestPlayerList bool `json:"requestPlayerList,omitempty"` - RequestNowPlaying bool `json:"requestNowPlaying,omitempty"` - RequestVolume bool `json:"requestVolume,omitempty"` - - // Action fields - Player string `json:"player,omitempty"` - Action string `json:"action,omitempty"` - SetVolume *int `json:"setVolume,omitempty"` - Seek *int64 `json:"Seek,omitempty"` - SetPosition *int64 `json:"SetPosition,omitempty"` - SetShuffle *bool `json:"setShuffle,omitempty"` - SetLoopStatus string `json:"setLoopStatus,omitempty"` - - // Album art - AlbumArtUrl string `json:"albumArtUrl,omitempty"` - - // Set when this packet carries album art bytes over a side channel. - TransferringAlbumArt bool `json:"transferringAlbumArt,omitempty"` - - // Player list from remote device - PlayerList []string `json:"playerList,omitempty"` - - // NowPlaying fields — populated when phone sends state update - Title string `json:"title,omitempty"` - Artist string `json:"artist,omitempty"` - Album string `json:"album,omitempty"` - Url string `json:"url,omitempty"` - Length int64 `json:"length,omitempty"` - Pos int64 `json:"pos,omitempty"` - IsPlaying bool `json:"isPlaying,omitempty"` - Volume int `json:"volume,omitempty"` - CanControl bool `json:"canControl,omitempty"` - CanGoNext bool `json:"canGoNext,omitempty"` - CanGoPrevious bool `json:"canGoPrevious,omitempty"` - CanPause bool `json:"canPause,omitempty"` - CanPlay bool `json:"canPlay,omitempty"` - CanSeek bool `json:"canSeek,omitempty"` - PlaybackStatus string `json:"playbackStatus,omitempty"` - Shuffle *bool `json:"shuffle,omitempty"` - LoopStatus string `json:"loopStatus,omitempty"` -} - -type NowPlaying struct { - Player string `json:"player"` - Title string `json:"title"` - Artist string `json:"artist"` - Album string `json:"album"` - AlbumArtUrl string `json:"albumArtUrl"` - ArtPending bool `json:"artPending,omitempty"` - Url string `json:"url,omitempty"` - Length int64 `json:"length"` - Pos int64 `json:"pos,omitempty"` - // PosAnchorMs is the wall-clock time (Unix millis) at which Pos was - // sampled. Clients compute the live position drift-free as - // Pos + (nowMs - PosAnchorMs) * (isPlaying ? 1 : 0). - PosAnchorMs int64 `json:"posAnchorMs,omitempty"` - IsPlaying bool `json:"isPlaying"` - Volume int `json:"volume,omitempty"` - CanControl bool `json:"canControl"` - CanGoNext bool `json:"canGoNext"` - CanGoPrevious bool `json:"canGoPrevious"` - CanPause bool `json:"canPause"` - CanPlay bool `json:"canPlay"` - CanSeek bool `json:"canSeek"` - PlaybackStatus string `json:"playbackStatus"` - Shuffle *bool `json:"shuffle,omitempty"` - LoopStatus string `json:"loopStatus,omitempty"` -} - -// DeepCopy returns a fully independent copy of NowPlaying. -// Pointer fields (Shuffle) are deep-copied to prevent shared-memory races -// between the cached state and callers. -func (p *NowPlaying) DeepCopy() *NowPlaying { - if p == nil { - return nil - } - cp := *p - if p.Shuffle != nil { - s := *p.Shuffle - cp.Shuffle = &s - } - return &cp -} - -func (p *MPRISPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - p.mu.Lock() - if _, exists := p.devices[dev.ID()]; !exists { - p.devices[dev.ID()] = dev - } - p.mu.Unlock() - - var body MPRISRequest - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return err - } - - if body.AlbumArtUrl != "" && strings.HasPrefix(body.AlbumArtUrl, "file://") { - go p.sendAlbumArt(ctx, dev, body.Player, body.AlbumArtUrl) - return nil - } - - // Inbound album art payload — the phone responds to requestAlbumArt - // with a side-channel transfer carrying the art bytes. - if body.TransferringAlbumArt && pkt.PayloadSize > 0 && pkt.PayloadTransferInfo != nil { - go p.receiveAlbumArt(ctx, dev, body.Player, body.AlbumArtUrl, - pkt.PayloadSize, pkt.PayloadTransferInfo.Port) - return nil - } - - if body.RequestPlayerList { - return p.sendPlayerList(dev) - } - - // Incoming playerList from remote device — prune players that no longer - // exist (their media session was destroyed) and request fresh status for - // the ones still around. - if body.PlayerList != nil { - p.logger.Debug("mpris: received player list from remote", zap.Strings("players", body.PlayerList)) - pruned := false - p.mu.Lock() - if prev := p.remoteStates[dev.ID()]; prev != nil { - inList := false - for _, name := range body.PlayerList { - if name == prev.Player { - inList = true - break - } - } - if !inList { - delete(p.remoteStates, dev.ID()) - delete(p.remoteStateTimes, dev.ID()) - delete(p.positionTrackers, dev.ID()) - pruned = true - } - } - p.mu.Unlock() - - if pruned && p.bus != nil { - // The tracked player's session is gone — emit an empty update so - // watchers fall back to "no media playing" for this device. - p.bus.Publish(events.TypeMprisUpdate, dev.ID(), &NowPlaying{}) - } - - for _, player := range body.PlayerList { - player := player - go p.requestPlayerStatus(dev, player) - } - return nil - } - - if body.Player == "" { - return nil - } - - if body.RequestNowPlaying || body.RequestVolume { - go func() { - if state, err := p.playerState(body.Player); err == nil { - p.broadcast(state) - } - }() - return nil - } - - if body.Action != "" || body.Seek != nil || body.SetPosition != nil || body.SetVolume != nil || body.SetShuffle != nil || body.SetLoopStatus != "" { - go p.handleAction(body.Player, body.Action, body.Seek, body.SetPosition, body.SetVolume, body.SetShuffle, body.SetLoopStatus) - return nil - } - - // Incoming NowPlaying state update from a remote device - if body.Title != "" || body.Artist != "" || body.Album != "" || body.IsPlaying || body.PlaybackStatus != "" { - state := &NowPlaying{ - Player: body.Player, - Title: body.Title, - Artist: body.Artist, - Album: body.Album, - AlbumArtUrl: body.AlbumArtUrl, - Url: body.Url, - Length: body.Length, - Pos: body.Pos, - IsPlaying: body.IsPlaying, - Volume: body.Volume, - CanControl: body.CanControl, - CanGoNext: body.CanGoNext, - CanGoPrevious: body.CanGoPrevious, - CanPause: body.CanPause, - CanPlay: body.CanPlay, - CanSeek: body.CanSeek, - PlaybackStatus: body.PlaybackStatus, - Shuffle: body.Shuffle, - LoopStatus: body.LoopStatus, - } - p.mu.Lock() - tracker, ok := p.positionTrackers[dev.ID()] - if !ok { - tracker = &remotePositionTracker{} - p.positionTrackers[dev.ID()] = tracker - } - tracker.lastPosition = body.Pos - tracker.lastPositionAt = time.Now() - tracker.playing = body.IsPlaying - state.PosAnchorMs = tracker.lastPositionAt.UnixMilli() - shouldPublish := shouldPublishRemoteState(p.remoteStates[dev.ID()], state) - p.remoteStates[dev.ID()] = state - p.remoteStateTimes[dev.ID()] = time.Now() - p.mu.Unlock() - - // Request art bytes from the phone when the advertsed album art is - // a kdeconnect:// URI we have not cached yet. Resolve any already - // cached art before publishing so watch clients get a loadable URL. - if p.artCache != nil && p.artCache.Resolve(state.AlbumArtUrl) == "" { - go p.requestAlbumArt(dev, state.Player, state.AlbumArtUrl) - } - if shouldPublish && p.bus != nil { - pub := state.DeepCopy() - if p.artCache != nil { - if resolved := p.artCache.Resolve(pub.AlbumArtUrl); resolved != "" { - pub.AlbumArtUrl = resolved - } else if pub.AlbumArtUrl != "" && strings.HasPrefix(pub.AlbumArtUrl, "kdeconnect:") { - // Art still downloading: publish an empty URL with the - // pending flag instead of an unloadable kdeconnect:/ - // URI. The arrival re-publish carries file://. - pub.AlbumArtUrl = "" - pub.ArtPending = true - } - } - p.bus.Publish(events.TypeMprisUpdate, dev.ID(), pub) - } - return nil - } - - return nil -} - -func shouldPublishRemoteState(last, current *NowPlaying) bool { - if last == nil { - return true - } - return last.Player != current.Player || - last.Title != current.Title || - last.Artist != current.Artist || - last.Album != current.Album || - last.AlbumArtUrl != current.AlbumArtUrl || - last.PlaybackStatus != current.PlaybackStatus || - last.IsPlaying != current.IsPlaying || - last.Volume != current.Volume -} - -func (p *MPRISPlugin) sendPlayerList(dev device.Sender) error { - p.mu.RLock() - displayNames := make([]string, 0, len(p.players)) - for name := range p.players { - displayNames = append(displayNames, name) - } - p.mu.RUnlock() - - if displayNames == nil { - displayNames = []string{} - } - - p.logger.Debug("mpris: sending player list", zap.Strings("players", displayNames)) - - pkt, err := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ - "playerList": displayNames, - "supportAlbumArtPayload": true, - }) - if err != nil { - return err - } - - go func() { - for _, name := range displayNames { - if state, err := p.playerState(name); err == nil { - p.broadcast(state) - } - } - }() - - return dev.Send(pkt) -} - -func (p *MPRISPlugin) sendPlayerListBroadcast() { - p.mu.RLock() - displayNames := make([]string, 0, len(p.players)) - for name := range p.players { - displayNames = append(displayNames, name) - } - p.mu.RUnlock() - - if displayNames == nil { - displayNames = []string{} - } - - pkt, _ := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ - "playerList": displayNames, - "supportAlbumArtPayload": true, - }) - - p.mu.RLock() - defer p.mu.RUnlock() - for _, dev := range p.devices { - if dev.IsConnected() { - _ = dev.Send(pkt) - } - } -} - -func (p *MPRISPlugin) broadcast(state *NowPlaying) { - pkt, err := protocol.NewPacket("kdeconnect.mpris", state) - if err != nil { - return - } - - p.mu.RLock() - defer p.mu.RUnlock() - for _, dev := range p.devices { - if dev.IsConnected() { - _ = dev.Send(pkt) - } - } -} - -func (p *MPRISPlugin) addPlayer(busName, uniqueName, displayName, shortName string) { - p.mu.Lock() - p.players[displayName] = &trackedPlayer{ - busName: busName, - uniqueName: uniqueName, - displayName: displayName, - shortName: shortName, - } - p.mu.Unlock() - - p.logger.Debug("mpris: added player", zap.String("displayName", displayName), zap.String("busName", busName)) - - if state, err := p.playerState(displayName); err == nil { - p.mu.Lock() - p.lastStates[displayName] = state - p.mu.Unlock() - p.broadcast(state) - } - - p.sendPlayerListBroadcast() -} - -func (p *MPRISPlugin) removePlayer(displayName string) { - p.mu.Lock() - delete(p.players, displayName) - delete(p.lastTracks, displayName) - delete(p.lastStates, displayName) - p.mu.Unlock() - - p.logger.Debug("mpris: removed player", zap.String("displayName", displayName)) - - p.sendPlayerListBroadcast() -} - -func (p *MPRISPlugin) resolvePlayer(displayName string) *trackedPlayer { - if displayName == "" { - return nil - } - p.mu.RLock() - defer p.mu.RUnlock() - for _, pl := range p.players { - if pl.displayName == displayName || pl.shortName == displayName || strings.EqualFold(pl.shortName, displayName) { - return pl - } - } - return nil -} - -func (p *MPRISPlugin) sendAlbumArt(ctx context.Context, dev device.Sender, player, artUrl string) { - p.mu.Lock() - reqKey := dev.ID() + "|" + artUrl - if lastReq, exists := p.artRequests[reqKey]; exists && time.Since(lastReq) < 5*time.Second { - p.mu.Unlock() - return - } - p.artRequests[reqKey] = time.Now() - p.mu.Unlock() - - cleanUrl := artUrl - if idx := strings.LastIndex(cleanUrl, "?t="); idx != -1 { - cleanUrl = cleanUrl[:idx] - } - - filePath := strings.TrimPrefix(cleanUrl, "file://") - if unescaped, err := url.PathUnescape(filePath); err == nil { - filePath = unescaped - } - - f, err := os.Open(filePath) - if err != nil { - p.logger.Debug("mpris: album art file not found", zap.String("path", filePath), zap.Error(err)) - return - } - stat, err := f.Stat() - f.Close() - if err != nil { - return - } - - var shareCfg config.ShareConfig - shareCfg.Defaults() - - ln, port, err := share.ListenSideChannel(ctx, shareCfg, p.tlsConfig) - if err != nil { - return - } - - go func() { - _ = share.AcceptAndSend(ln, filePath, p.tlsConfig, dev.ID(), cert.PinnedFingerprint(dev.PeerCert()), 10*time.Second, nil, p.logger) - }() - - pkt, err := protocol.NewPacket("kdeconnect.mpris", map[string]interface{}{ - "transferringAlbumArt": true, - "player": player, - "albumArtUrl": artUrl, - }) - if err == nil { - pkt.PayloadSize = stat.Size() - pkt.PayloadTransferInfo = &protocol.TransferInfo{Port: port} - dev.Send(pkt) - } -} - -// requestAlbumArt asks the remote device to stream the album art bytes -// referenced by a kdeconnect:/artUri URI over a side channel. -func (p *MPRISPlugin) requestAlbumArt(dev device.Sender, player, artUrl string) { - if player == "" || artUrl == "" { - return - } - p.mu.Lock() - reqKey := "req|" + dev.ID() + "|" + artUrl - if lastReq, exists := p.artRequests[reqKey]; exists && time.Since(lastReq) < 10*time.Second { - p.mu.Unlock() - return - } - p.artRequests[reqKey] = time.Now() - p.mu.Unlock() - - body := MPRISRequest{ - Player: player, - AlbumArtUrl: artUrl, - } - pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) - if err != nil { - return - } - if err := dev.Send(pkt); err != nil { - p.logger.Debug("mpris: album art request failed", zap.Error(err)) - } -} - -// receiveAlbumArt streams an inbound album art payload into the cache -// and re-publishes the device state with a resolved file:// URL so watch -// clients and kcd mpris status surface the loadable location. -func (p *MPRISPlugin) receiveAlbumArt(_ context.Context, dev device.Sender, player, artUrl string, size int64, port int) { - remoteIP := dev.RemoteIP() - if remoteIP == nil { - return - } - if p.artCache == nil || size <= 0 || size > maxAlbumArtBytes { - return - } - - p.logger.Debug("mpris: receiving album art from remote", - zap.String("player", player), - zap.String("device_id", dev.ID()), - zap.Int64("size", size)) - - tmp, err := os.CreateTemp(p.artCache.Dir(), ".art-*") - if err != nil { - p.logger.Warn("mpris: failed to create temp file for album art", zap.Error(err)) - return - } - tmpPath := tmp.Name() - tmp.Close() - defer os.Remove(tmpPath) - - // The Handle ctx is canceled as soon as Handle returns; use an - // independent context so the side-channel dial isn't aborted. - dlCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - if err := share.ReceiveSideChannel(dlCtx, remoteIP, port, size, tmpPath, p.tlsConfig, cert.PinnedFingerprint(dev.PeerCert()), nil, p.logger); err != nil { - p.logger.Warn("mpris: album art transfer failed", zap.Error(err)) - return - } - - fileURL, err := p.artCache.Commit(artUrl, tmpPath) - if err != nil { - p.logger.Warn("mpris: failed to cache album art", zap.Error(err)) - return - } - - p.stampAlbumArt(dev.ID(), player, artUrl, fileURL) -} - -// stampAlbumArt publishes a resolved file:// URL for a completed album-art -// fetch — but only if the device's current track still wants exactly this -// art. The side-channel download can take up to 30s, during which the track -// (or the active player) may have changed; stamping blindly would show the -// old track's cover on the new track. On mismatch the bytes stay in the art -// cache and the new track's own art request fulfills it. -func (p *MPRISPlugin) stampAlbumArt(deviceID, player, artUrl, fileURL string) { - p.mu.Lock() - state := p.remoteStates[deviceID] - if state == nil || state.AlbumArtUrl != artUrl || state.Player != player { - p.mu.Unlock() - return - } - state = state.DeepCopy() - state.AlbumArtUrl = fileURL - p.mu.Unlock() - if p.bus != nil { - p.bus.Publish(events.TypeMprisUpdate, deviceID, state) - } -} - -func (p *MPRISPlugin) OnConnect(dev device.Sender) { - p.logger.Info("mpris: device connected, requesting player list", zap.String("device_id", dev.ID())) - go p.requestPlayerListPeriodic(dev) -} - -func (p *MPRISPlugin) requestPlayerListPeriodic(dev device.Sender) { - if !dev.IsConnected() { - return - } - p.requestPlayerList(dev) - timer := time.NewTimer(3 * time.Second) - <-timer.C - if !dev.IsConnected() { - return - } - p.requestPlayerList(dev) - timer.Reset(7 * time.Second) - <-timer.C - if !dev.IsConnected() { - return - } - p.requestPlayerList(dev) -} - -// remoteStatePollInterval is how often the daemon re-requests now-playing -// from devices with an active remote player. Clients are then pure-push: -// fresh state arrives within one interval of connect, and position stays -// current without any client-side polling. -const remoteStatePollInterval = 5 * time.Second - -// startRemoteStatePoller periodically re-requests now-playing from every -// connected device that has a known active player. The responses flow back -// through Handle, where shouldPublishRemoteState dedupes them, so an -// mpris.update is only republished when the state actually changes — not -// on every poll. This closes the "watch client misses mid-track state" -// gap from the initial dump's 10s freshness gate. -func (p *MPRISPlugin) startRemoteStatePoller(ctx context.Context) { - go func() { - ticker := time.NewTicker(remoteStatePollInterval) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - p.pollRemoteStates() - } - } - }() -} - -// pollRemoteStates requests a now-playing refresh from devices that have a -// cached, actively-playing player. Devices without a cached state (never -// reported a player) or whose player is stopped/paused are skipped — stopped -// players are intentionally left to go stale instead of keeping a ghost track -// perpetually fresh. -func (p *MPRISPlugin) pollRemoteStates() { - p.mu.RLock() - type target struct { - dev device.Sender - player string - } - var targets []target - for id, dev := range p.devices { - if !dev.IsConnected() { - continue - } - state := p.remoteStates[id] - if state == nil || state.Player == "" || !state.IsPlaying { - continue - } - targets = append(targets, target{dev: dev, player: state.Player}) - } - p.mu.RUnlock() - - for _, t := range targets { - if err := p.requestPlayerStatus(t.dev, t.player); err != nil { - p.logger.Debug("mpris: state poll request failed", zap.Error(err)) - } - } -} - -func (p *MPRISPlugin) OnDisconnect(dev device.Sender) { - p.mu.Lock() - defer p.mu.Unlock() - delete(p.devices, dev.ID()) - delete(p.remoteStates, dev.ID()) - delete(p.remoteStateTimes, dev.ID()) - delete(p.positionTrackers, dev.ID()) -} - -type DebugPlayerInfo struct { - DisplayName string `json:"displayName"` - BusName string `json:"busName"` - ShortName string `json:"shortName"` - Title string `json:"title"` - Artist string `json:"artist"` - Album string `json:"album"` - PlaybackStatus string `json:"playbackStatus"` - IsPlaying bool `json:"isPlaying"` - Volume int `json:"volume"` - Pos int64 `json:"pos"` - Length int64 `json:"length"` - AlbumArtUrl string `json:"albumArtUrl"` - CanSeek bool `json:"canSeek"` - CanGoNext bool `json:"canGoNext"` - CanGoPrevious bool `json:"canGoPrevious"` - CanPlay bool `json:"canPlay"` - CanPause bool `json:"canPause"` - Error string `json:"error,omitempty"` -} - -type DebugStatus struct { - WatcherRunning bool `json:"watcherRunning"` - DeviceCount int `json:"deviceCount"` - Players []DebugPlayerInfo `json:"players"` - PlayerMappings map[string]string `json:"playerMappings"` -} - -func (p *MPRISPlugin) DebugStatus() *DebugStatus { - p.mu.RLock() - watching := p.watching - devCount := len(p.devices) - playerMappings := make(map[string]string, len(p.players)) - playerList := make([]*trackedPlayer, 0, len(p.players)) - for _, pl := range p.players { - playerMappings[pl.displayName] = pl.busName - playerList = append(playerList, pl) - } - p.mu.RUnlock() - - var players []DebugPlayerInfo - for _, pl := range playerList { - info := DebugPlayerInfo{ - DisplayName: pl.displayName, - BusName: pl.busName, - ShortName: pl.shortName, - } - if state, err := p.playerState(pl.displayName); err == nil { - info.Title = state.Title - info.Artist = state.Artist - info.Album = state.Album - info.PlaybackStatus = state.PlaybackStatus - info.IsPlaying = state.IsPlaying - info.Volume = state.Volume - info.Pos = state.Pos - info.Length = state.Length - info.AlbumArtUrl = state.AlbumArtUrl - info.CanSeek = state.CanSeek - info.CanGoNext = state.CanGoNext - info.CanGoPrevious = state.CanGoPrevious - info.CanPlay = state.CanPlay - info.CanPause = state.CanPause - } else { - info.Error = err.Error() - } - players = append(players, info) - } - - return &DebugStatus{ - WatcherRunning: watching, - DeviceCount: devCount, - Players: players, - PlayerMappings: playerMappings, - } -} - -// SendAction sends a media control action to a remote device. -// Sends on both kdeconnect.mpris (for Android's old MprisPlugin) and -// kdeconnect.mpris.request (for MprisReceiverPlugin) to maximise compatibility. -func (p *MPRISPlugin) SendAction(dev device.Sender, player, action string, seek *int64, volume *int) error { - body := MPRISRequest{ - Player: player, - Action: action, - SetVolume: volume, - Seek: seek, - } - pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) - if err != nil { - return err - } - if err := dev.Send(pkt); err != nil { - return err - } - pkt2, err := protocol.NewPacket("kdeconnect.mpris", body) - if err != nil { - return err - } - return dev.Send(pkt2) -} - -// requestPlayerList sends a request for the remote device's active player list. -func (p *MPRISPlugin) requestPlayerList(dev device.Sender) error { - body := MPRISRequest{RequestPlayerList: true} - pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) - if err != nil { - return err - } - if err := dev.Send(pkt); err != nil { - return err - } - // Also send as kdeconnect.mpris for phone-side MprisReceiverPlugin/MprisPlugin. - pkt2, err := protocol.NewPacket("kdeconnect.mpris", body) - if err != nil { - return err - } - return dev.Send(pkt2) -} - -// requestPlayerStatus sends a requestNowPlaying + requestVolume for a specific player. -func (p *MPRISPlugin) requestPlayerStatus(dev device.Sender, player string) error { - body := MPRISRequest{ - Player: player, - RequestNowPlaying: true, - RequestVolume: true, - } - pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// RequestState sends a requestNowPlaying to refresh remote state. -func (p *MPRISPlugin) RequestState(dev device.Sender, player string) error { - if player == "" { - return p.requestPlayerList(dev) - } - return p.requestPlayerStatus(dev, player) -} - -// RemoteState returns the last known NowPlaying state for a remote device, -// markArtPending empties unloadable kdeconnect:/ art URIs and flags them, -// so serving paths (status, snapshot, summaries) agree with published -// events: art is either a loadable URL or "" with ArtPending set. -func markArtPending(np *NowPlaying) { - if np.AlbumArtUrl != "" && strings.HasPrefix(np.AlbumArtUrl, "kdeconnect:") { - np.AlbumArtUrl = "" - np.ArtPending = true - } -} - -// with position extrapolated from the last update time if playing. -func (p *MPRISPlugin) RemoteState(deviceID string) *NowPlaying { - p.mu.RLock() - defer p.mu.RUnlock() - state := p.remoteStates[deviceID] - if state == nil { - return nil - } - copy := state.DeepCopy() - copy.AlbumArtUrl = p.resolveArtURL(copy.AlbumArtUrl) - markArtPending(copy) - if tracker, ok := p.positionTrackers[deviceID]; ok && tracker.playing { - elapsed := time.Since(tracker.lastPositionAt).Milliseconds() - copy.Pos = tracker.lastPosition + elapsed - } - return copy -} - -// RemoteStateAge returns the time since the last remote state update for a device. -// Returns a large duration if no state has been received yet. -func (p *MPRISPlugin) RemoteStateAge(deviceID string) time.Duration { - p.mu.RLock() - defer p.mu.RUnlock() - t, ok := p.remoteStateTimes[deviceID] - if !ok { - return 365 * 24 * time.Hour // effectively "forever ago" - } - return time.Since(t) -} - -// RemoteStates returns all known remote device player states, -// with positions extrapolated from the last update time. -func (p *MPRISPlugin) RemoteStates() map[string]*NowPlaying { - p.mu.RLock() - defer p.mu.RUnlock() - result := make(map[string]*NowPlaying, len(p.remoteStates)) - for id, state := range p.remoteStates { - if state == nil { - continue - } - copy := state.DeepCopy() - copy.AlbumArtUrl = p.resolveArtURL(copy.AlbumArtUrl) - markArtPending(copy) - if tracker, ok := p.positionTrackers[id]; ok && tracker.playing { - elapsed := time.Since(tracker.lastPositionAt).Milliseconds() - copy.Pos = tracker.lastPosition + elapsed - } - result[id] = copy - } - return result -} - -// resolveArtURL maps a cached kdeconnect:// album art URI to a loadable -// file:// URL. Non-kdeconnect URIs and not-yet-cached art pass through. -func (p *MPRISPlugin) resolveArtURL(raw string) string { - if raw == "" || p.artCache == nil { - return raw - } - if resolved := p.artCache.Resolve(raw); resolved != "" { - return resolved - } - return raw -} - -// watchTelephony subscribes to telephony events and pauses/resumes -// local MPRIS players when calls start/end. -func (p *MPRISPlugin) watchTelephony(ctx context.Context) { - sub := p.bus.Subscribe(events.DefaultSubscriberCap, - events.TypeTelephonyRinging, - events.TypeTelephonyTalking, - events.TypeTelephonyCanceled) - defer sub.Close() - - for { - select { - case <-ctx.Done(): - return - case ev := <-sub.C: - p.handleTelephonyEvent(ev) - } - } -} - -func (p *MPRISPlugin) handleTelephonyEvent(ev events.Event) { - switch ev.Type { - case events.TypeTelephonyRinging, events.TypeTelephonyTalking: - p.pauseAllPlayers() - case events.TypeTelephonyCanceled: - p.resumePausedPlayers() - } -} - -// pauseAllPlayers pauses every currently-playing local MPRIS player. -// It only acts once per call — repeated ringing/talking events are no-ops -// while callPausedPlayers is non-empty. -func (p *MPRISPlugin) pauseAllPlayers() { - p.mu.Lock() - defer p.mu.Unlock() - - if len(p.callPausedPlayers) > 0 { - return // already paused for an active call - } - - for name, pl := range p.players { - state, err := p.playerStateDBus(pl.busName, name) - if err != nil { - continue - } - if !state.IsPlaying { - continue - } - obj := p.dbus.Object(pl.busName, "/org/mpris/MediaPlayer2") - if err := dbusCall(obj, "org.mpris.MediaPlayer2.Player.Pause").Err; err == nil { - p.callPausedPlayers = append(p.callPausedPlayers, name) - } - } -} - -// resumePausedPlayers resumes every player that was paused by pauseAllPlayers. -func (p *MPRISPlugin) resumePausedPlayers() { - p.mu.Lock() - defer p.mu.Unlock() - - for _, name := range p.callPausedPlayers { - pl := p.players[name] - if pl == nil { - continue - } - obj := p.dbus.Object(pl.busName, "/org/mpris/MediaPlayer2") - _ = dbusCall(obj, "org.mpris.MediaPlayer2.Player.Play").Err - } - p.callPausedPlayers = p.callPausedPlayers[:0] -} - -// ActivePlayers returns the list of player names from all remote device states. -func (p *MPRISPlugin) ActivePlayers() []string { - p.mu.RLock() - defer p.mu.RUnlock() - seen := make(map[string]struct{}) - for _, state := range p.remoteStates { - if state != nil && state.Player != "" { - seen[state.Player] = struct{}{} - } - } - players := make([]string, 0, len(seen)) - for name := range seen { - players = append(players, name) - } - return players -} diff --git a/internal/plugins/mpris/remote.go b/internal/plugins/mpris/remote.go new file mode 100644 index 0000000..3e27c83 --- /dev/null +++ b/internal/plugins/mpris/remote.go @@ -0,0 +1,252 @@ +package mpris + +import ( + "context" + "strings" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// SendAction sends a media control action to a remote device. +// Sends on both kdeconnect.mpris (for Android's old MprisPlugin) and +// kdeconnect.mpris.request (for MprisReceiverPlugin) to maximise compatibility. +func (p *MPRISPlugin) SendAction(dev device.Sender, player, action string, seek *int64, volume *int) error { + body := MPRISRequest{ + Player: player, + Action: action, + SetVolume: volume, + Seek: seek, + } + pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) + if err != nil { + return err + } + if err := dev.Send(pkt); err != nil { + return err + } + pkt2, err := protocol.NewPacket("kdeconnect.mpris", body) + if err != nil { + return err + } + return dev.Send(pkt2) +} + +// requestPlayerList sends a request for the remote device's active player list. +func (p *MPRISPlugin) requestPlayerList(dev device.Sender) error { + body := MPRISRequest{RequestPlayerList: true} + pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) + if err != nil { + return err + } + if err := dev.Send(pkt); err != nil { + return err + } + // Also send as kdeconnect.mpris for phone-side MprisReceiverPlugin/MprisPlugin. + pkt2, err := protocol.NewPacket("kdeconnect.mpris", body) + if err != nil { + return err + } + return dev.Send(pkt2) +} + +// requestPlayerStatus sends a requestNowPlaying + requestVolume for a specific player. +func (p *MPRISPlugin) requestPlayerStatus(dev device.Sender, player string) error { + body := MPRISRequest{ + Player: player, + RequestNowPlaying: true, + RequestVolume: true, + } + pkt, err := protocol.NewPacket("kdeconnect.mpris.request", body) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// RequestState sends a requestNowPlaying to refresh remote state. +func (p *MPRISPlugin) RequestState(dev device.Sender, player string) error { + if player == "" { + return p.requestPlayerList(dev) + } + return p.requestPlayerStatus(dev, player) +} + +// markArtPending empties unloadable kdeconnect:/ art URIs and flags them, +// so serving paths (status, snapshot, summaries) agree with published +// events: art is either a loadable URL or "" with ArtPending set. +func markArtPending(np *NowPlaying) { + if np.AlbumArtUrl != "" && strings.HasPrefix(np.AlbumArtUrl, "kdeconnect:") { + np.AlbumArtUrl = "" + np.ArtPending = true + } +} + +// RemoteState returns the last known NowPlaying state for a remote device, +// with position extrapolated from the last update time if playing. +func (p *MPRISPlugin) RemoteState(deviceID string) *NowPlaying { + p.mu.RLock() + defer p.mu.RUnlock() + state := p.remoteStates[deviceID] + if state == nil { + return nil + } + copy := state.DeepCopy() + copy.AlbumArtUrl = p.resolveArtURL(copy.AlbumArtUrl) + markArtPending(copy) + if tracker, ok := p.positionTrackers[deviceID]; ok && tracker.playing { + elapsed := time.Since(tracker.lastPositionAt).Milliseconds() + copy.Pos = tracker.lastPosition + elapsed + } + return copy +} + +// RemoteStateAge returns the time since the last remote state update for a device. +// Returns a large duration if no state has been received yet. +func (p *MPRISPlugin) RemoteStateAge(deviceID string) time.Duration { + p.mu.RLock() + defer p.mu.RUnlock() + t, ok := p.remoteStateTimes[deviceID] + if !ok { + return 365 * 24 * time.Hour // effectively "forever ago" + } + return time.Since(t) +} + +// RemoteStates returns all known remote device player states, +// with positions extrapolated from the last update time. +func (p *MPRISPlugin) RemoteStates() map[string]*NowPlaying { + p.mu.RLock() + defer p.mu.RUnlock() + result := make(map[string]*NowPlaying, len(p.remoteStates)) + for id, state := range p.remoteStates { + if state == nil { + continue + } + copy := state.DeepCopy() + copy.AlbumArtUrl = p.resolveArtURL(copy.AlbumArtUrl) + markArtPending(copy) + if tracker, ok := p.positionTrackers[id]; ok && tracker.playing { + elapsed := time.Since(tracker.lastPositionAt).Milliseconds() + copy.Pos = tracker.lastPosition + elapsed + } + result[id] = copy + } + return result +} + +// resolveArtURL maps a cached kdeconnect:// album art URI to a loadable +// file:// URL. Non-kdeconnect URIs and not-yet-cached art pass through. +func (p *MPRISPlugin) resolveArtURL(raw string) string { + if raw == "" || p.artCache == nil { + return raw + } + if resolved := p.artCache.Resolve(raw); resolved != "" { + return resolved + } + return raw +} + +// ActivePlayers returns the list of player names from all remote device states. +func (p *MPRISPlugin) ActivePlayers() []string { + p.mu.RLock() + defer p.mu.RUnlock() + seen := make(map[string]struct{}) + for _, state := range p.remoteStates { + if state != nil && state.Player != "" { + seen[state.Player] = struct{}{} + } + } + players := make([]string, 0, len(seen)) + for name := range seen { + players = append(players, name) + } + return players +} + +func (p *MPRISPlugin) OnConnect(dev device.Sender) { + p.logger.Info("mpris: device connected, requesting player list", zap.String("device_id", dev.ID())) + go p.requestPlayerListPeriodic(dev) +} + +func (p *MPRISPlugin) requestPlayerListPeriodic(dev device.Sender) { + if !dev.IsConnected() { + return + } + p.requestPlayerList(dev) + timer := time.NewTimer(3 * time.Second) + <-timer.C + if !dev.IsConnected() { + return + } + p.requestPlayerList(dev) + timer.Reset(7 * time.Second) + <-timer.C + if !dev.IsConnected() { + return + } + p.requestPlayerList(dev) +} + +// startRemoteStatePoller periodically re-requests now-playing from every +// connected device that has a known active player. The responses flow back +// through Handle, where shouldPublishRemoteState dedupes them, so an +// mpris.update is only republished when the state actually changes — not +// on every poll. This closes the "watch client misses mid-track state" +// gap from the initial dump's 10s freshness gate. +func (p *MPRISPlugin) startRemoteStatePoller(ctx context.Context) { + go func() { + ticker := time.NewTicker(remoteStatePollInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + p.pollRemoteStates() + } + } + }() +} + +// pollRemoteStates requests a now-playing refresh from devices that have a +// cached, actively-playing player. Devices without a cached state (never +// reported a player) or whose player is stopped/paused are skipped — stopped +// players are intentionally left to go stale instead of keeping a ghost track +// perpetually fresh. +func (p *MPRISPlugin) pollRemoteStates() { + p.mu.RLock() + type target struct { + dev device.Sender + player string + } + var targets []target + for id, dev := range p.devices { + if !dev.IsConnected() { + continue + } + state := p.remoteStates[id] + if state == nil || state.Player == "" || !state.IsPlaying { + continue + } + targets = append(targets, target{dev: dev, player: state.Player}) + } + p.mu.RUnlock() + + for _, t := range targets { + if err := p.requestPlayerStatus(t.dev, t.player); err != nil { + p.logger.Debug("mpris: state poll request failed", zap.Error(err)) + } + } +} + +func (p *MPRISPlugin) OnDisconnect(dev device.Sender) { + p.mu.Lock() + defer p.mu.Unlock() + delete(p.devices, dev.ID()) + delete(p.remoteStates, dev.ID()) + delete(p.remoteStateTimes, dev.ID()) + delete(p.positionTrackers, dev.ID()) +} diff --git a/internal/plugins/mpris/telephony.go b/internal/plugins/mpris/telephony.go new file mode 100644 index 0000000..2dafc9a --- /dev/null +++ b/internal/plugins/mpris/telephony.go @@ -0,0 +1,77 @@ +package mpris + +import ( + "context" + + "github.com/bethropolis/kcd/internal/events" +) + +// watchTelephony subscribes to telephony events and pauses/resumes +// local MPRIS players when calls start/end. +func (p *MPRISPlugin) watchTelephony(ctx context.Context) { + sub := p.bus.Subscribe(events.DefaultSubscriberCap, + events.TypeTelephonyRinging, + events.TypeTelephonyTalking, + events.TypeTelephonyCanceled) + defer sub.Close() + + for { + select { + case <-ctx.Done(): + return + case ev := <-sub.C: + p.handleTelephonyEvent(ev) + } + } +} + +func (p *MPRISPlugin) handleTelephonyEvent(ev events.Event) { + switch ev.Type { + case events.TypeTelephonyRinging, events.TypeTelephonyTalking: + p.pauseAllPlayers() + case events.TypeTelephonyCanceled: + p.resumePausedPlayers() + } +} + +// pauseAllPlayers pauses every currently-playing local MPRIS player. +// It only acts once per call — repeated ringing/talking events are no-ops +// while callPausedPlayers is non-empty. +func (p *MPRISPlugin) pauseAllPlayers() { + p.mu.Lock() + defer p.mu.Unlock() + + if len(p.callPausedPlayers) > 0 { + return // already paused for an active call + } + + for name, pl := range p.players { + state, err := p.playerStateDBus(pl.busName, name) + if err != nil { + continue + } + if !state.IsPlaying { + continue + } + obj := p.dbus.Object(pl.busName, "/org/mpris/MediaPlayer2") + if err := dbusCall(obj, "org.mpris.MediaPlayer2.Player.Pause").Err; err == nil { + p.callPausedPlayers = append(p.callPausedPlayers, name) + } + } +} + +// resumePausedPlayers resumes every player that was paused by pauseAllPlayers. +func (p *MPRISPlugin) resumePausedPlayers() { + p.mu.Lock() + defer p.mu.Unlock() + + for _, name := range p.callPausedPlayers { + pl := p.players[name] + if pl == nil { + continue + } + obj := p.dbus.Object(pl.busName, "/org/mpris/MediaPlayer2") + _ = dbusCall(obj, "org.mpris.MediaPlayer2.Player.Play").Err + } + p.callPausedPlayers = p.callPausedPlayers[:0] +} diff --git a/internal/plugins/mpris/types.go b/internal/plugins/mpris/types.go new file mode 100644 index 0000000..b59f188 --- /dev/null +++ b/internal/plugins/mpris/types.go @@ -0,0 +1,123 @@ +package mpris + +import "time" + +type MPRISRequest struct { + // Request fields + RequestPlayerList bool `json:"requestPlayerList,omitempty"` + RequestNowPlaying bool `json:"requestNowPlaying,omitempty"` + RequestVolume bool `json:"requestVolume,omitempty"` + + // Action fields + Player string `json:"player,omitempty"` + Action string `json:"action,omitempty"` + SetVolume *int `json:"setVolume,omitempty"` + Seek *int64 `json:"Seek,omitempty"` + SetPosition *int64 `json:"SetPosition,omitempty"` + SetShuffle *bool `json:"setShuffle,omitempty"` + SetLoopStatus string `json:"setLoopStatus,omitempty"` + + // Album art + AlbumArtUrl string `json:"albumArtUrl,omitempty"` + + // Set when this packet carries album art bytes over a side channel. + TransferringAlbumArt bool `json:"transferringAlbumArt,omitempty"` + + // Player list from remote device + PlayerList []string `json:"playerList,omitempty"` + + // NowPlaying fields — populated when phone sends state update + Title string `json:"title,omitempty"` + Artist string `json:"artist,omitempty"` + Album string `json:"album,omitempty"` + Url string `json:"url,omitempty"` + Length int64 `json:"length,omitempty"` + Pos int64 `json:"pos,omitempty"` + IsPlaying bool `json:"isPlaying,omitempty"` + Volume int `json:"volume,omitempty"` + CanControl bool `json:"canControl,omitempty"` + CanGoNext bool `json:"canGoNext,omitempty"` + CanGoPrevious bool `json:"canGoPrevious,omitempty"` + CanPause bool `json:"canPause,omitempty"` + CanPlay bool `json:"canPlay,omitempty"` + CanSeek bool `json:"canSeek,omitempty"` + PlaybackStatus string `json:"playbackStatus,omitempty"` + Shuffle *bool `json:"shuffle,omitempty"` + LoopStatus string `json:"loopStatus,omitempty"` +} + +type NowPlaying struct { + Player string `json:"player"` + Title string `json:"title"` + Artist string `json:"artist"` + Album string `json:"album"` + AlbumArtUrl string `json:"albumArtUrl"` + ArtPending bool `json:"artPending,omitempty"` + Url string `json:"url,omitempty"` + Length int64 `json:"length"` + Pos int64 `json:"pos,omitempty"` + // PosAnchorMs is the wall-clock time (Unix millis) at which Pos was + // sampled. Clients compute the live position drift-free as + // Pos + (nowMs - PosAnchorMs) * (isPlaying ? 1 : 0). + PosAnchorMs int64 `json:"posAnchorMs,omitempty"` + IsPlaying bool `json:"isPlaying"` + Volume int `json:"volume,omitempty"` + CanControl bool `json:"canControl"` + CanGoNext bool `json:"canGoNext"` + CanGoPrevious bool `json:"canGoPrevious"` + CanPause bool `json:"canPause"` + CanPlay bool `json:"canPlay"` + CanSeek bool `json:"canSeek"` + PlaybackStatus string `json:"playbackStatus"` + Shuffle *bool `json:"shuffle,omitempty"` + LoopStatus string `json:"loopStatus,omitempty"` +} + +// DeepCopy returns a fully independent copy of NowPlaying. +// Pointer fields (Shuffle) are deep-copied to prevent shared-memory races +// between the cached state and callers. +func (p *NowPlaying) DeepCopy() *NowPlaying { + if p == nil { + return nil + } + cp := *p + if p.Shuffle != nil { + s := *p.Shuffle + cp.Shuffle = &s + } + return &cp +} + +type DebugPlayerInfo struct { + DisplayName string `json:"displayName"` + BusName string `json:"busName"` + ShortName string `json:"shortName"` + Title string `json:"title"` + Artist string `json:"artist"` + Album string `json:"album"` + PlaybackStatus string `json:"playbackStatus"` + IsPlaying bool `json:"isPlaying"` + Volume int `json:"volume"` + Pos int64 `json:"pos"` + Length int64 `json:"length"` + AlbumArtUrl string `json:"albumArtUrl"` + CanSeek bool `json:"canSeek"` + CanGoNext bool `json:"canGoNext"` + CanGoPrevious bool `json:"canGoPrevious"` + CanPlay bool `json:"canPlay"` + CanPause bool `json:"canPause"` + Error string `json:"error,omitempty"` +} + +type DebugStatus struct { + WatcherRunning bool `json:"watcherRunning"` + DeviceCount int `json:"deviceCount"` + Players []DebugPlayerInfo `json:"players"` + PlayerMappings map[string]string `json:"playerMappings"` +} + +// remoteStatePollInterval is how often the daemon re-requests now-playing +// from devices with an active remote player. Clients are then pure-push: +// fresh state arrives within one interval of connect, and position stays +// current without any client-side polling. +const remoteStatePollInterval = 5 * time.Second diff --git a/internal/plugins/notification/actions.go b/internal/plugins/notification/actions.go new file mode 100644 index 0000000..73f79f7 --- /dev/null +++ b/internal/plugins/notification/actions.go @@ -0,0 +1,43 @@ +package notification + +import ( + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" +) + +// RequestReply sends a reply back to an Android notification. +func (p *NotificationPlugin) RequestReply(dev device.Sender, replyID, message string) error { + pkt, err := protocol.NewPacket("kdeconnect.notification.reply", map[string]string{ + "requestReplyId": replyID, + "message": message, + }) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// Dismiss asks the phone to clear the notification with the given ID and +// closes the matching desktop popup, if one is tracked. +func (p *NotificationPlugin) Dismiss(dev device.Sender, id string) error { + pkt, err := protocol.NewPacket("kdeconnect.notification.request", map[string]string{ + "cancel": id, + }) + if err != nil { + return err + } + if err := dev.Send(pkt); err != nil { + return err + } + if id != "" { + if desktopID, ok := p.notifIDs.LoadAndDelete(p.notifKey(dev.ID(), id)); ok { + if s, ok := desktopID.(string); ok { + p.closeNotification(s) + } + } + } + return nil +} + +func (p *NotificationPlugin) OnConnect(_ device.Sender) {} +func (p *NotificationPlugin) OnDisconnect(_ device.Sender) {} diff --git a/internal/plugins/notification/desktop.go b/internal/plugins/notification/desktop.go new file mode 100644 index 0000000..ffec2dc --- /dev/null +++ b/internal/plugins/notification/desktop.go @@ -0,0 +1,86 @@ +package notification + +import ( + "context" + "strconv" + "strings" + "time" +) + +// closeNotification closes a previously-shown desktop popup by id. +func (p *NotificationPlugin) closeNotification(desktopID string) { + go func() { + _ = p.newExec(context.Background(), "gdbus", "call", "--session", + "--dest", "org.freedesktop.Notifications", + "--object-path", "/org/freedesktop/Notifications", + "--method", "org.freedesktop.Notifications.CloseNotification", + desktopID, + ).Run() + }() +} + +// sendDesktopNotification calls notify-send with the collected parameters. +func (p *NotificationPlugin) sendDesktopNotification(devID, appName, id, title, text, iconPath string) { + // Dunst / mako / swaync: stack notifications from the same app so they + // replace each other instead of flooding the screen. Daemons that ignore + // this hint (e.g. Quickshell) are covered by --replace-id below. + groupHint := "string:x-dunst-stack-tag:kcd-" + appName + + args := []string{"-a", appName} + if p.cfg.Urgency != "" { + args = append(args, "-u", p.cfg.Urgency) + } + if p.cfg.ExpireMS >= 0 { + args = append(args, "-t", strconv.Itoa(p.cfg.ExpireMS)) + } + + // No icon by default. Pass an explicit empty icon so daemons (e.g. + // Quickshell) don't fall back to deriving an icon name from the app name + // and render a placeholder. When show_icons is enabled, pass the phone's + // downloaded icon, falling back to a name derived from the app. + iconArg := "" + if p.cfg.ShowIcons { + iconArg = strings.ToLower(strings.ReplaceAll(appName, " ", "-")) + if iconPath != "" { + iconArg = iconPath + } + if iconArg == "" { + iconArg = "smartphone" + } + } + args = append(args, "-i", iconArg, "-h", groupHint) + + if p.canCloseNotifs && id != "" { + // Android re-posts notifications on every update with a stable id + // (e.g. media/scrobble now-playing). Replace the existing desktop + // popup via D-Bus replaces_id so repeated updates collapse to one, + // mirroring the reference desktop's Notification::update(). Disable + // via `replace_notifications = false`. + if p.cfg.ReplaceNotifications { + key := p.notifKey(devID, id) + // Cancel-grace: if a cancel was deferred for this key, drop it so + // the popup is updated in place rather than closed and re-opened. + if t, ok := p.pendingCloses.LoadAndDelete(key); ok { + t.(*time.Timer).Stop() + } + if prevID, ok := p.notifIDs.Load(key); ok { + if s, ok := prevID.(string); ok && s != "" { + args = append(args, "-r", s) + } + } + } + // "--" ends option parsing so a phone-provided title/body starting with + // '-' can't be misparsed as a notify-send flag (notify-send uses + // GOption, which honors the POSIX separator). + args = append(args, "--print-id", "--", title, text) + } else { + args = append(args, "--", title, text) + } + + out, err := p.newExec(context.Background(), "notify-send", args...).Output() + if err == nil && p.canCloseNotifs && id != "" { + if desktopID := strings.TrimSpace(string(out)); desktopID != "" { + p.notifIDs.Store(p.notifKey(devID, id), desktopID) + } + } +} diff --git a/internal/plugins/notification/handle.go b/internal/plugins/notification/handle.go new file mode 100644 index 0000000..06e32f8 --- /dev/null +++ b/internal/plugins/notification/handle.go @@ -0,0 +1,139 @@ +package notification + +import ( + "context" + "encoding/json" + "net" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" +) + +// Handle processes an incoming notification. +func (p *NotificationPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + var body NotificationBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return err + } + + // Handle cancellation — close the corresponding desktop notification. + if body.IsCancel { + if body.ID != "" { + key := p.notifKey(dev.ID(), body.ID) + if desktopID, ok := p.notifIDs.Load(key); ok { + if p.cfg.CancelGraceMS > 0 { + // Cancel-grace: hold the popup open so a same-id re-post + // (media now-playing toggling play/pause) updates it in + // place instead of closing and re-opening. If no re-post + // arrives, close after the grace window. + if t, ok := p.pendingCloses.Load(key); ok { + t.(*time.Timer).Stop() + } + want := desktopID.(string) + p.pendingCloses.Store(key, time.AfterFunc( + time.Duration(p.cfg.CancelGraceMS)*time.Millisecond, + func() { + if cur, ok := p.notifIDs.LoadAndDelete(key); ok { + // Only close the popup we scheduled — if a + // re-post replaced it, leave the new one. + if cur.(string) == want { + p.closeNotification(want) + } + } + p.pendingCloses.Delete(key) + }, + )) + } else { + p.notifIDs.Delete(key) + p.closeNotification(desktopID.(string)) + } + } + } + if p.bus != nil { + p.bus.Publish(events.TypeNotificationCanceled, dev.ID(), map[string]string{"id": body.ID}) + } + return nil + } + + if body.Silent { + return nil + } + + // Skip non-clearable notifications (media playback, foreground services). + if p.cfg.SkipNonClearable && !body.IsClearable { + return nil + } + + // Apply per-app notification filter. + action := p.resolveAction(body.AppName) + if action == "silent" { + // Still publish the event for scripts/watch, but skip the desktop popup. + if p.bus != nil { + payload := map[string]any{ + "appName": body.AppName, + "title": body.Title, + "text": body.Text, + } + if body.RequestReplyId != "" { + payload["requestReplyId"] = body.RequestReplyId + } + p.bus.Publish(events.TypeNotification, dev.ID(), payload) + } + return nil + } + + // Truncate text to keep notifications readable. + text := body.Text + if p.cfg.MaxBodyLength > 0 && len(text) > p.cfg.MaxBodyLength { + text = text[:p.cfg.MaxBodyLength] + "…" + } else if p.cfg.MaxBodyLength == 0 && len(text) > 512 { + // Maintain the old default limit if no config is set + text = text[:512] + "…" + } + + appName := nonAlphaNumeric.ReplaceAllString(body.AppName, "") + + if p.bus != nil { + payload := map[string]any{ + "appName": body.AppName, + "title": body.Title, + "text": body.Text, + } + if body.RequestReplyId != "" { + payload["requestReplyId"] = body.RequestReplyId + } + p.bus.Publish(events.TypeNotification, dev.ID(), payload) + } + + // Capture payload info before the goroutine — pkt may be released. + var ( + hasIcon = pkt.PayloadSize > 0 && pkt.PayloadTransferInfo != nil + payloadSize = pkt.PayloadSize + payloadPort int + remoteIP net.IP + expectedFP string + ) + if hasIcon { + payloadPort = pkt.PayloadTransferInfo.Port + remoteIP = dev.RemoteIP() + if remoteIP == nil { + hasIcon = false + } else { + expectedFP = cert.PinnedFingerprint(dev.PeerCert()) + } + } + + // Handlers must not block — all I/O in a goroutine. + go func() { + var iconPath string + if p.cfg.ShowIcons { + iconPath = p.fetchIcon(ctx, appName, body.ID, remoteIP, payloadPort, payloadSize, hasIcon, expectedFP) + } + p.sendDesktopNotification(dev.ID(), appName, body.ID, body.Title, text, iconPath) + }() + + return nil +} diff --git a/internal/plugins/notification/icon.go b/internal/plugins/notification/icon.go new file mode 100644 index 0000000..a3a4300 --- /dev/null +++ b/internal/plugins/notification/icon.go @@ -0,0 +1,101 @@ +package notification + +import ( + "context" + "fmt" + "io" + "net" + "os" + "path/filepath" + "regexp" + "strings" + + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// notifIDChars restricts phone-provided notification IDs to filename-safe +// characters when they're embedded in icon cache paths. +var notifIDChars = regexp.MustCompile(`[^a-zA-Z0-9._-]`) + +// sanitizeNotifID strips everything but alphanumerics, dot, underscore and +// hyphen from a notification ID so it can't escape the icon cache directory +// via path separators. Empty results fall back to "default". +func sanitizeNotifID(id string) string { + if safe := notifIDChars.ReplaceAllString(id, "_"); safe != "" { + return safe + } + return "default" +} + +// notifKey scopes a notification id to its device so two paired phones with +// colliding Android notification keys don't replace each other's popups. +func (p *NotificationPlugin) notifKey(devID, id string) string { + return devID + "|" + id +} + +// fetchIcon downloads the notification icon payload and returns the path to the +// saved file, or an empty string if unavailable. +func (p *NotificationPlugin) fetchIcon( + ctx context.Context, + appName, notifID string, + remoteIP net.IP, + port int, + size int64, + hasIcon bool, + expectedFP string, +) string { + if !p.cfg.FetchIcons || p.tlsConfig == nil || p.iconDir == "" { + // Fall back to icon name derived from app name. + return "" + } + + // Use the notification ID as the filename so the same app reuses the + // cached icon rather than downloading it on every notification. + // The ID comes from the phone, so restrict it to filename-safe chars + // (no separators) — otherwise ../ in an ID would escape the icon dir. + safeName := nonAlphaNumeric.ReplaceAllString(appName, "_") + safeID := sanitizeNotifID(notifID) + iconPath := filepath.Join(p.iconDir, fmt.Sprintf("%s-%s.png", safeName, safeID)) + + // Belt and braces: confine the result to the icon dir even if the + // sanitizer above ever regresses. + if rel, err := filepath.Rel(p.iconDir, iconPath); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { + p.logger.Warn("notification: icon path escapes cache dir, refusing", + zap.String("id", notifID)) + return "" + } + + // Reuse the cached icon even when the phone re-posts the notification + // without an icon payload (Android only sends the bytes when the icon + // hash changes). Without this, every re-post would fall back to a theme + // icon name that doesn't exist, showing a placeholder image. + if _, err := os.Stat(iconPath); err == nil { + return iconPath + } + + if !hasIcon { + return "" + } + + conn, err := transport.DialSidechannel(ctx, remoteIP, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) + if err != nil { + p.logger.Debug("notification: icon dial failed", zap.Error(err)) + return "" + } + defer conn.Close() + + f, err := os.Create(iconPath) + if err != nil { + return "" + } + defer f.Close() + + if _, err := io.Copy(f, io.LimitReader(conn, size)); err != nil { + p.logger.Debug("notification: icon download failed", zap.Error(err)) + _ = os.Remove(iconPath) + return "" + } + + return iconPath +} diff --git a/internal/plugins/notification/notification.go b/internal/plugins/notification/notification.go deleted file mode 100644 index f9d416b..0000000 --- a/internal/plugins/notification/notification.go +++ /dev/null @@ -1,463 +0,0 @@ -package notification - -import ( - "context" - "crypto/tls" - "encoding/json" - "fmt" - "io" - "net" - "os" - "os/exec" - "path/filepath" - "regexp" - "strconv" - "strings" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "github.com/bethropolis/kcd/internal/transport" - "go.uber.org/zap" -) - -// NotificationPlugin handles incoming notifications and displays them on the desktop. -type NotificationPlugin struct { - sidechannel transport.SidechannelOptions - bus *events.Bus - tlsConfig *tls.Config - logger *zap.Logger - notifIDs sync.Map // maps deviceID|body.ID -> desktop notify-send ID (string) - pendingCloses sync.Map // maps deviceID|body.ID -> *time.Timer (deferred close for cancel-grace) - iconDir string // temp dir for cached notification icons - cfg config.NotificationPluginConfig - canCloseNotifs bool // whether notify-send supports --print-id - mu sync.RWMutex - filters config.NotificationConfig - newExec func(ctx context.Context, name string, args ...string) *exec.Cmd -} - -// NewNotificationPlugin creates a NotificationPlugin. -// tlsConfig is used to fetch notification icon payloads over the KDE Connect -// side-channel; pass nil to disable icon fetching. -func NewNotificationPlugin(cfg config.NotificationPluginConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...transport.SidechannelOptions) *NotificationPlugin { - var sidechannel transport.SidechannelOptions - if len(options) > 0 { - sidechannel = options[0] - } - p := &NotificationPlugin{ - sidechannel: sidechannel, - cfg: cfg, - bus: bus, - tlsConfig: tlsConfig, - logger: logger.With(zap.String("plugin", "notification")), - newExec: exec.CommandContext, - } - - // Probe --print-id support by checking --help output. - // This is side-effect-free and immune to version string format changes. - if out, err := exec.CommandContext(context.Background(), "notify-send", "--help").CombinedOutput(); err == nil { - p.canCloseNotifs = strings.Contains(string(out), "--print-id") - } - - // Create a persistent temp directory for icon files so they survive - // long enough for the notification daemon to read them. - baseDir := cfg.IconCacheDir - if dir, err := os.MkdirTemp(baseDir, "kcd-notif-icons-*"); err == nil { - p.iconDir = dir - } - - return p -} - -// Close removes the icon temp directory. Call when the plugin is no longer needed. -func (p *NotificationPlugin) Close() { - if p.iconDir != "" { - _ = os.RemoveAll(p.iconDir) - } - p.pendingCloses.Range(func(k, v any) bool { - if t, ok := v.(*time.Timer); ok { - t.Stop() - } - return true - }) -} - -// SetFilters atomically replaces the per-app notification filter map. -func (p *NotificationPlugin) SetFilters(f config.NotificationConfig) { - p.mu.Lock() - p.filters = f - p.mu.Unlock() -} - -// resolveAction returns the configured action for an app ("show" or "silent"). -func (p *NotificationPlugin) resolveAction(appName string) string { - p.mu.RLock() - f := p.filters - p.mu.RUnlock() - if f == nil { - return "show" - } - if action, ok := f[appName]; ok { - return action - } - if def, ok := f["*"]; ok { - return def - } - return "show" -} - -// NotificationBody represents the fields of a notification packet. -type NotificationBody struct { - ID string `json:"id"` - AppName string `json:"appName"` - Title string `json:"title"` - Text string `json:"text"` - IsCancel bool `json:"isCancel,omitempty"` - IsClearable bool `json:"isClearable,omitempty"` - Silent bool `json:"silent,omitempty"` - RequestReplyId string `json:"requestReplyId,omitempty"` -} - -func (p *NotificationPlugin) Name() string { return "Notification" } -func (p *NotificationPlugin) Timeout() time.Duration { return 5 * time.Second } -func (p *NotificationPlugin) IncomingTypes() []string { - return []string{"kdeconnect.notification"} -} -func (p *NotificationPlugin) OutgoingTypes() []string { - return []string{"kdeconnect.notification.reply", "kdeconnect.notification.request"} -} - -// nonAlphaNumeric sanitises app names to be safe for exec / notify-send args. -var nonAlphaNumeric = regexp.MustCompile(`[^a-zA-Z0-9 ._-]`) - -// Handle processes an incoming notification. -func (p *NotificationPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - var body NotificationBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return err - } - - // Handle cancellation — close the corresponding desktop notification. - if body.IsCancel { - if body.ID != "" { - key := p.notifKey(dev.ID(), body.ID) - if desktopID, ok := p.notifIDs.Load(key); ok { - if p.cfg.CancelGraceMS > 0 { - // Cancel-grace: hold the popup open so a same-id re-post - // (media now-playing toggling play/pause) updates it in - // place instead of closing and re-opening. If no re-post - // arrives, close after the grace window. - if t, ok := p.pendingCloses.Load(key); ok { - t.(*time.Timer).Stop() - } - want := desktopID.(string) - p.pendingCloses.Store(key, time.AfterFunc( - time.Duration(p.cfg.CancelGraceMS)*time.Millisecond, - func() { - if cur, ok := p.notifIDs.LoadAndDelete(key); ok { - // Only close the popup we scheduled — if a - // re-post replaced it, leave the new one. - if cur.(string) == want { - p.closeNotification(want) - } - } - p.pendingCloses.Delete(key) - }, - )) - } else { - p.notifIDs.Delete(key) - p.closeNotification(desktopID.(string)) - } - } - } - if p.bus != nil { - p.bus.Publish(events.TypeNotificationCanceled, dev.ID(), map[string]string{"id": body.ID}) - } - return nil - } - - if body.Silent { - return nil - } - - // Skip non-clearable notifications (media playback, foreground services). - if p.cfg.SkipNonClearable && !body.IsClearable { - return nil - } - - // Apply per-app notification filter. - action := p.resolveAction(body.AppName) - if action == "silent" { - // Still publish the event for scripts/watch, but skip the desktop popup. - if p.bus != nil { - payload := map[string]any{ - "appName": body.AppName, - "title": body.Title, - "text": body.Text, - } - if body.RequestReplyId != "" { - payload["requestReplyId"] = body.RequestReplyId - } - p.bus.Publish(events.TypeNotification, dev.ID(), payload) - } - return nil - } - - // Truncate text to keep notifications readable. - text := body.Text - if p.cfg.MaxBodyLength > 0 && len(text) > p.cfg.MaxBodyLength { - text = text[:p.cfg.MaxBodyLength] + "…" - } else if p.cfg.MaxBodyLength == 0 && len(text) > 512 { - // Maintain the old default limit if no config is set - text = text[:512] + "…" - } - - appName := nonAlphaNumeric.ReplaceAllString(body.AppName, "") - - if p.bus != nil { - payload := map[string]any{ - "appName": body.AppName, - "title": body.Title, - "text": body.Text, - } - if body.RequestReplyId != "" { - payload["requestReplyId"] = body.RequestReplyId - } - p.bus.Publish(events.TypeNotification, dev.ID(), payload) - } - - // Capture payload info before the goroutine — pkt may be released. - var ( - hasIcon = pkt.PayloadSize > 0 && pkt.PayloadTransferInfo != nil - payloadSize = pkt.PayloadSize - payloadPort int - remoteIP net.IP - expectedFP string - ) - if hasIcon { - payloadPort = pkt.PayloadTransferInfo.Port - remoteIP = dev.RemoteIP() - if remoteIP == nil { - hasIcon = false - } else { - expectedFP = cert.PinnedFingerprint(dev.PeerCert()) - } - } - - // Handlers must not block — all I/O in a goroutine. - go func() { - var iconPath string - if p.cfg.ShowIcons { - iconPath = p.fetchIcon(ctx, appName, body.ID, remoteIP, payloadPort, payloadSize, hasIcon, expectedFP) - } - p.sendDesktopNotification(dev.ID(), appName, body.ID, body.Title, text, iconPath) - }() - - return nil -} - -// notifIDChars restricts phone-provided notification IDs to filename-safe -// characters when they're embedded in icon cache paths. -var notifIDChars = regexp.MustCompile(`[^a-zA-Z0-9._-]`) - -// sanitizeNotifID strips everything but alphanumerics, dot, underscore and -// hyphen from a notification ID so it can't escape the icon cache directory -// via path separators. Empty results fall back to "default". -func sanitizeNotifID(id string) string { - if safe := notifIDChars.ReplaceAllString(id, "_"); safe != "" { - return safe - } - return "default" -} - -// notifKey scopes a notification id to its device so two paired phones with -// colliding Android notification keys don't replace each other's popups. -func (p *NotificationPlugin) notifKey(devID, id string) string { - return devID + "|" + id -} - -// fetchIcon downloads the notification icon payload and returns the path to the -// saved file, or an empty string if unavailable. -func (p *NotificationPlugin) fetchIcon( - ctx context.Context, - appName, notifID string, - remoteIP net.IP, - port int, - size int64, - hasIcon bool, - expectedFP string, -) string { - if !p.cfg.FetchIcons || p.tlsConfig == nil || p.iconDir == "" { - // Fall back to icon name derived from app name. - return "" - } - - // Use the notification ID as the filename so the same app reuses the - // cached icon rather than downloading it on every notification. - // The ID comes from the phone, so restrict it to filename-safe chars - // (no separators) — otherwise ../ in an ID would escape the icon dir. - safeName := nonAlphaNumeric.ReplaceAllString(appName, "_") - safeID := sanitizeNotifID(notifID) - iconPath := filepath.Join(p.iconDir, fmt.Sprintf("%s-%s.png", safeName, safeID)) - - // Belt and braces: confine the result to the icon dir even if the - // sanitizer above ever regresses. - if rel, err := filepath.Rel(p.iconDir, iconPath); err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { - p.logger.Warn("notification: icon path escapes cache dir, refusing", - zap.String("id", notifID)) - return "" - } - - // Reuse the cached icon even when the phone re-posts the notification - // without an icon payload (Android only sends the bytes when the icon - // hash changes). Without this, every re-post would fall back to a theme - // icon name that doesn't exist, showing a placeholder image. - if _, err := os.Stat(iconPath); err == nil { - return iconPath - } - - if !hasIcon { - return "" - } - - conn, err := transport.DialSidechannel(ctx, remoteIP, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) - if err != nil { - p.logger.Debug("notification: icon dial failed", zap.Error(err)) - return "" - } - defer conn.Close() - - f, err := os.Create(iconPath) - if err != nil { - return "" - } - defer f.Close() - - if _, err := io.Copy(f, io.LimitReader(conn, size)); err != nil { - p.logger.Debug("notification: icon download failed", zap.Error(err)) - _ = os.Remove(iconPath) - return "" - } - - return iconPath -} - -// closeNotification closes a previously-shown desktop popup by id. -func (p *NotificationPlugin) closeNotification(desktopID string) { - go func() { - _ = p.newExec(context.Background(), "gdbus", "call", "--session", - "--dest", "org.freedesktop.Notifications", - "--object-path", "/org/freedesktop/Notifications", - "--method", "org.freedesktop.Notifications.CloseNotification", - desktopID, - ).Run() - }() -} - -// sendDesktopNotification calls notify-send with the collected parameters. -func (p *NotificationPlugin) sendDesktopNotification(devID, appName, id, title, text, iconPath string) { - // Dunst / mako / swaync: stack notifications from the same app so they - // replace each other instead of flooding the screen. Daemons that ignore - // this hint (e.g. Quickshell) are covered by --replace-id below. - groupHint := "string:x-dunst-stack-tag:kcd-" + appName - - args := []string{"-a", appName} - if p.cfg.Urgency != "" { - args = append(args, "-u", p.cfg.Urgency) - } - if p.cfg.ExpireMS >= 0 { - args = append(args, "-t", strconv.Itoa(p.cfg.ExpireMS)) - } - - // No icon by default. Pass an explicit empty icon so daemons (e.g. - // Quickshell) don't fall back to deriving an icon name from the app name - // and render a placeholder. When show_icons is enabled, pass the phone's - // downloaded icon, falling back to a name derived from the app. - iconArg := "" - if p.cfg.ShowIcons { - iconArg = strings.ToLower(strings.ReplaceAll(appName, " ", "-")) - if iconPath != "" { - iconArg = iconPath - } - if iconArg == "" { - iconArg = "smartphone" - } - } - args = append(args, "-i", iconArg, "-h", groupHint) - - if p.canCloseNotifs && id != "" { - // Android re-posts notifications on every update with a stable id - // (e.g. media/scrobble now-playing). Replace the existing desktop - // popup via D-Bus replaces_id so repeated updates collapse to one, - // mirroring the reference desktop's Notification::update(). Disable - // via `replace_notifications = false`. - if p.cfg.ReplaceNotifications { - key := p.notifKey(devID, id) - // Cancel-grace: if a cancel was deferred for this key, drop it so - // the popup is updated in place rather than closed and re-opened. - if t, ok := p.pendingCloses.LoadAndDelete(key); ok { - t.(*time.Timer).Stop() - } - if prevID, ok := p.notifIDs.Load(key); ok { - if s, ok := prevID.(string); ok && s != "" { - args = append(args, "-r", s) - } - } - } - // "--" ends option parsing so a phone-provided title/body starting with - // '-' can't be misparsed as a notify-send flag (notify-send uses - // GOption, which honors the POSIX separator). - args = append(args, "--print-id", "--", title, text) - } else { - args = append(args, "--", title, text) - } - - out, err := p.newExec(context.Background(), "notify-send", args...).Output() - if err == nil && p.canCloseNotifs && id != "" { - if desktopID := strings.TrimSpace(string(out)); desktopID != "" { - p.notifIDs.Store(p.notifKey(devID, id), desktopID) - } - } -} - -// RequestReply sends a reply back to an Android notification. -func (p *NotificationPlugin) RequestReply(dev device.Sender, replyID, message string) error { - pkt, err := protocol.NewPacket("kdeconnect.notification.reply", map[string]string{ - "requestReplyId": replyID, - "message": message, - }) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// Dismiss asks the phone to clear the notification with the given ID and -// closes the matching desktop popup, if one is tracked. -func (p *NotificationPlugin) Dismiss(dev device.Sender, id string) error { - pkt, err := protocol.NewPacket("kdeconnect.notification.request", map[string]string{ - "cancel": id, - }) - if err != nil { - return err - } - if err := dev.Send(pkt); err != nil { - return err - } - if id != "" { - if desktopID, ok := p.notifIDs.LoadAndDelete(p.notifKey(dev.ID(), id)); ok { - if s, ok := desktopID.(string); ok { - p.closeNotification(s) - } - } - } - return nil -} - -func (p *NotificationPlugin) OnConnect(_ device.Sender) {} -func (p *NotificationPlugin) OnDisconnect(_ device.Sender) {} diff --git a/internal/plugins/notification/types.go b/internal/plugins/notification/types.go new file mode 100644 index 0000000..6bb8168 --- /dev/null +++ b/internal/plugins/notification/types.go @@ -0,0 +1,127 @@ +package notification + +import ( + "context" + "crypto/tls" + "os" + "os/exec" + "regexp" + "strings" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// NotificationPlugin handles incoming notifications and displays them on the desktop. +type NotificationPlugin struct { + sidechannel transport.SidechannelOptions + bus *events.Bus + tlsConfig *tls.Config + logger *zap.Logger + notifIDs sync.Map // maps deviceID|body.ID -> desktop notify-send ID (string) + pendingCloses sync.Map // maps deviceID|body.ID -> *time.Timer (deferred close for cancel-grace) + iconDir string // temp dir for cached notification icons + cfg config.NotificationPluginConfig + canCloseNotifs bool // whether notify-send supports --print-id + mu sync.RWMutex + filters config.NotificationConfig + newExec func(ctx context.Context, name string, args ...string) *exec.Cmd +} + +// NewNotificationPlugin creates a NotificationPlugin. +// tlsConfig is used to fetch notification icon payloads over the KDE Connect +// side-channel; pass nil to disable icon fetching. +func NewNotificationPlugin(cfg config.NotificationPluginConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...transport.SidechannelOptions) *NotificationPlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } + p := &NotificationPlugin{ + sidechannel: sidechannel, + cfg: cfg, + bus: bus, + tlsConfig: tlsConfig, + logger: logger.With(zap.String("plugin", "notification")), + newExec: exec.CommandContext, + } + + // Probe --print-id support by checking --help output. + // This is side-effect-free and immune to version string format changes. + if out, err := exec.CommandContext(context.Background(), "notify-send", "--help").CombinedOutput(); err == nil { + p.canCloseNotifs = strings.Contains(string(out), "--print-id") + } + + // Create a persistent temp directory for icon files so they survive + // long enough for the notification daemon to read them. + baseDir := cfg.IconCacheDir + if dir, err := os.MkdirTemp(baseDir, "kcd-notif-icons-*"); err == nil { + p.iconDir = dir + } + + return p +} + +// Close removes the icon temp directory. Call when the plugin is no longer needed. +func (p *NotificationPlugin) Close() { + if p.iconDir != "" { + _ = os.RemoveAll(p.iconDir) + } + p.pendingCloses.Range(func(k, v any) bool { + if t, ok := v.(*time.Timer); ok { + t.Stop() + } + return true + }) +} + +// SetFilters atomically replaces the per-app notification filter map. +func (p *NotificationPlugin) SetFilters(f config.NotificationConfig) { + p.mu.Lock() + p.filters = f + p.mu.Unlock() +} + +// resolveAction returns the configured action for an app ("show" or "silent"). +func (p *NotificationPlugin) resolveAction(appName string) string { + p.mu.RLock() + f := p.filters + p.mu.RUnlock() + if f == nil { + return "show" + } + if action, ok := f[appName]; ok { + return action + } + if def, ok := f["*"]; ok { + return def + } + return "show" +} + +// NotificationBody represents the fields of a notification packet. +type NotificationBody struct { + ID string `json:"id"` + AppName string `json:"appName"` + Title string `json:"title"` + Text string `json:"text"` + IsCancel bool `json:"isCancel,omitempty"` + IsClearable bool `json:"isClearable,omitempty"` + Silent bool `json:"silent,omitempty"` + RequestReplyId string `json:"requestReplyId,omitempty"` +} + +func (p *NotificationPlugin) Name() string { return "Notification" } +func (p *NotificationPlugin) Timeout() time.Duration { return 5 * time.Second } +func (p *NotificationPlugin) IncomingTypes() []string { + return []string{"kdeconnect.notification"} +} +func (p *NotificationPlugin) OutgoingTypes() []string { + return []string{"kdeconnect.notification.reply", "kdeconnect.notification.request"} +} + +// nonAlphaNumeric sanitises app names to be safe for exec / notify-send args. +var nonAlphaNumeric = regexp.MustCompile(`[^a-zA-Z0-9 ._-]`) diff --git a/internal/plugins/pair/actions.go b/internal/plugins/pair/actions.go new file mode 100644 index 0000000..740b142 --- /dev/null +++ b/internal/plugins/pair/actions.go @@ -0,0 +1,141 @@ +package pair + +import ( + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// AcceptPairing accepts an incoming pair request. +func (p *PairPlugin) AcceptPairing(dev *device.Device) error { + pkt, err := protocol.NewPairPacket(protocol.PairAccept, 0) + if err != nil { + return err + } + + if err := dev.Send(pkt); err != nil { + p.logger.Error("failed to send pair accept", zap.Error(err)) + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + return err + } + + p.pairingDone(dev) + return nil +} + +// RequestPairing initiates a pairing request to a device. +func (p *PairPlugin) RequestPairing(dev *device.Device) error { + if dev.State() == device.StatePaired { + p.logger.Warn("device already paired", zap.String("device_id", dev.ID())) + return nil + } + + if dev.State() == device.StatePairRequestedByPeer { + // They already requested, just accept + return p.AcceptPairing(dev) + } + + // The request timestamp seeds the verification code on both sides, + // so generate it once and send exactly what we store. + timestamp := time.Now().Unix() + pkt, err := protocol.NewPairPacket(protocol.PairAccept, timestamp) + if err != nil { + return err + } + + p.mu.Lock() + p.pairingTimestamp[dev.ID()] = timestamp + p.mu.Unlock() + + if err := dev.Send(pkt); err != nil { + p.logger.Error("failed to send pair request", zap.Error(err)) + return err + } + + peerCert := dev.PeerCert() + if peerCert != nil { + vKey := cert.VerificationKey(p.localCert, peerCert, timestamp) + p.logger.Info("pairing verification code", + zap.String("device_id", dev.ID()), + zap.String("code", vKey)) + } + + dev.SetState(device.StatePairRequested) + p.logger.Info("pair request sent", zap.String("device_id", dev.ID())) + + if p.onStateChanged != nil { + p.onStateChanged() + } + + return nil +} + +// RejectPairing rejects an incoming pair request. +func (p *PairPlugin) RejectPairing(dev *device.Device) error { + pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) + if err != nil { + return err + } + + dev.Send(pkt) // best effort + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + + p.mu.Lock() + delete(p.pairingTimestamp, dev.ID()) + p.mu.Unlock() + + if p.onStateChanged != nil { + p.onStateChanged() + } + + p.logger.Info("pair request rejected", zap.String("device_id", dev.ID())) + return nil +} + +// Unpair removes pairing with a device. +func (p *PairPlugin) Unpair(dev *device.Device) error { + pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) + if err != nil { + return err + } + + dev.Send(pkt) // best effort + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + + p.mu.Lock() + delete(p.pairingTimestamp, dev.ID()) + p.mu.Unlock() + + if p.onStateChanged != nil { + p.onStateChanged() + } + + p.logger.Info("device unpaired", zap.String("device_id", dev.ID())) + return nil +} + +func (p *PairPlugin) pairingDone(dev *device.Device) { + dev.SetState(device.StatePaired) + dev.ClearPairDial() + + p.mu.Lock() + delete(p.pairingTimestamp, dev.ID()) + p.mu.Unlock() + + if p.onStateChanged != nil { + p.onStateChanged() + } + + p.logger.Info("pairing complete", zap.String("device_id", dev.ID())) + p.emit(events.TypePairAccepted, dev, "") +} diff --git a/internal/plugins/pair/handle.go b/internal/plugins/pair/handle.go new file mode 100644 index 0000000..019fce8 --- /dev/null +++ b/internal/plugins/pair/handle.go @@ -0,0 +1,142 @@ +package pair + +import ( + "context" + "encoding/json" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *PairPlugin) Handle(ctx context.Context, sender device.Sender, pkt *protocol.Packet) error { + var body protocol.PairBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return err + } + + dev, ok := p.devices.Get(sender.ID()) + if !ok { + p.logger.Warn("pair packet from unknown device", zap.String("device_id", sender.ID())) + return nil + } + + if body.Pair { + return p.handlePairRequest(ctx, dev, body) + } + return p.handleUnpairRequest(ctx, dev) +} + +func (p *PairPlugin) handlePairRequest(_ context.Context, dev *device.Device, body protocol.PairBody) error { + state := dev.State() + + switch state { + case device.StatePairRequested: + // We requested pairing, they accepted + p.logger.Info("pairing accepted by peer", zap.String("device_id", dev.ID())) + p.pairingDone(dev) + + case device.StatePairRequestedByPeer: + // Already have a pending request, ignore duplicate + p.logger.Debug("ignoring duplicate pair request", zap.String("device_id", dev.ID())) + + case device.StatePaired: + // Already paired - this is normal behavior in KDE Connect. + // The peer sends pair:true as confirmation/keep-alive. + // Just acknowledge by sending pair:true back. + p.logger.Debug("received pair confirmation from already paired device", zap.String("device_id", dev.ID())) + pkt, _ := protocol.NewPairPacket(protocol.PairAccept, 0) + dev.Send(pkt) + return nil + + case device.StateUnpaired, device.StateUnknown: + // New pair request from peer + // Validate timestamp for protocol v8 + if body.Timestamp > 0 { + now := time.Now().Unix() + diff := now - body.Timestamp + if diff < -AllowedTimestampDiff || diff > AllowedTimestampDiff { + p.logger.Warn("pair request timestamp out of range", + zap.String("device_id", dev.ID()), + zap.Int64("timestamp", body.Timestamp), + zap.Int64("now", now)) + // Send rejection + pkt, _ := protocol.NewPairPacket(protocol.PairReject, 0) + dev.Send(pkt) + return nil + } + // Store timestamp for verification key + p.mu.Lock() + p.pairingTimestamp[dev.ID()] = body.Timestamp + p.mu.Unlock() + } + + p.logger.Info("incoming pair request", zap.String("device_id", dev.ID())) + + var vKey string + peerCert := dev.PeerCert() + if peerCert != nil { + vKey = cert.VerificationKey(p.localCert, peerCert, p.pairingTimestampFor(dev.ID())) + p.logger.Info("pairing verification code", + zap.String("device_id", dev.ID()), + zap.String("code", vKey)) + } + + // Set state and wait for user to accept via CLI + dev.SetState(device.StatePairRequestedByPeer) + if p.onStateChanged != nil { + p.onStateChanged() + } + + p.emit(events.TypePairRequested, dev, vKey) + } + + return nil +} + +func (p *PairPlugin) handleUnpairRequest(_ context.Context, dev *device.Device) error { + state := dev.State() + + switch state { + case device.StatePairRequested: + // We requested, they rejected + p.logger.Info("pair request rejected by peer", zap.String("device_id", dev.ID())) + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + p.emit(events.TypePairRejected, dev, "") + + case device.StatePairRequestedByPeer: + // They requested, then cancelled + p.logger.Info("pair request cancelled by peer", zap.String("device_id", dev.ID())) + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + p.emit(events.TypePairRejected, dev, "") + + case device.StatePaired: + // Unpair request + p.logger.Info("unpair request received", zap.String("device_id", dev.ID())) + dev.SetState(device.StateUnpaired) + dev.ClearEphemeral() + dev.ClearPairDial() + + case device.StateUnpaired, device.StateUnknown: + // Already unpaired, ignore + p.logger.Debug("ignoring unpair request for unpaired device", zap.String("device_id", dev.ID())) + } + + // Clean up stored timestamp + p.mu.Lock() + delete(p.pairingTimestamp, dev.ID()) + p.mu.Unlock() + + if p.onStateChanged != nil { + p.onStateChanged() + } + + return nil +} diff --git a/internal/plugins/pair/pair.go b/internal/plugins/pair/pair.go deleted file mode 100644 index b2c034d..0000000 --- a/internal/plugins/pair/pair.go +++ /dev/null @@ -1,345 +0,0 @@ -package pair - -import ( - "context" - "crypto/x509" - "encoding/json" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "go.uber.org/zap" -) - -const ( - // AllowedTimestampDiff is the maximum allowed time difference for pairing timestamps (30 min) - AllowedTimestampDiff = 1800 -) - -// PairPlugin handles KDE Connect pairing protocol. -type PairPlugin struct { - devices *device.Registry - localCert *x509.Certificate - onStateChanged func() // callback to persist state - logger *zap.Logger - bus *events.Bus - cfg config.PairingConfig - - mu sync.Mutex - pairingTimestamp map[string]int64 // deviceID -> timestamp from pair request -} - -// NewPairPlugin creates a new pairing plugin. -func NewPairPlugin(devices *device.Registry, localCert *x509.Certificate, cfg config.PairingConfig, onStateChanged func(), bus *events.Bus, logger *zap.Logger) *PairPlugin { - return &PairPlugin{ - devices: devices, - localCert: localCert, - cfg: cfg, - onStateChanged: onStateChanged, - logger: logger.Named("pair"), - bus: bus, - pairingTimestamp: make(map[string]int64), - } -} - -// pairingTimestampFor returns the stored pair-request timestamp for a -// device, or zero if none was recorded (pre-v8 peer or unknown). -func (p *PairPlugin) pairingTimestampFor(deviceID string) int64 { - p.mu.Lock() - defer p.mu.Unlock() - return p.pairingTimestamp[deviceID] -} - -// emit publishes an event to the bus if one is configured. -func (p *PairPlugin) emit(typ events.EventType, dev *device.Device, vKey string) { - if p.bus == nil { - return - } - payload := map[string]interface{}{ - "name": dev.Name(), - "type": dev.Type, - } - if vKey != "" { - payload["verificationKey"] = vKey - } - p.bus.Publish(typ, dev.ID(), payload) -} - -func (p *PairPlugin) Name() string { return "Pair" } - -func (p *PairPlugin) Timeout() time.Duration { return 5 * time.Second } - -func (p *PairPlugin) IncomingTypes() []string { - return []string{protocol.TypePair} -} - -func (p *PairPlugin) OutgoingTypes() []string { - return []string{protocol.TypePair} -} - -func (p *PairPlugin) OnConnect(dev device.Sender) {} - -func (p *PairPlugin) OnDisconnect(dev device.Sender) {} - -func (p *PairPlugin) Handle(ctx context.Context, sender device.Sender, pkt *protocol.Packet) error { - var body protocol.PairBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return err - } - - dev, ok := p.devices.Get(sender.ID()) - if !ok { - p.logger.Warn("pair packet from unknown device", zap.String("device_id", sender.ID())) - return nil - } - - if body.Pair { - return p.handlePairRequest(ctx, dev, body) - } - return p.handleUnpairRequest(ctx, dev) -} - -func (p *PairPlugin) handlePairRequest(_ context.Context, dev *device.Device, body protocol.PairBody) error { - state := dev.State() - - switch state { - case device.StatePairRequested: - // We requested pairing, they accepted - p.logger.Info("pairing accepted by peer", zap.String("device_id", dev.ID())) - p.pairingDone(dev) - - case device.StatePairRequestedByPeer: - // Already have a pending request, ignore duplicate - p.logger.Debug("ignoring duplicate pair request", zap.String("device_id", dev.ID())) - - case device.StatePaired: - // Already paired - this is normal behavior in KDE Connect. - // The peer sends pair:true as confirmation/keep-alive. - // Just acknowledge by sending pair:true back. - p.logger.Debug("received pair confirmation from already paired device", zap.String("device_id", dev.ID())) - pkt, _ := protocol.NewPairPacket(protocol.PairAccept, 0) - dev.Send(pkt) - return nil - - case device.StateUnpaired, device.StateUnknown: - // New pair request from peer - // Validate timestamp for protocol v8 - if body.Timestamp > 0 { - now := time.Now().Unix() - diff := now - body.Timestamp - if diff < -AllowedTimestampDiff || diff > AllowedTimestampDiff { - p.logger.Warn("pair request timestamp out of range", - zap.String("device_id", dev.ID()), - zap.Int64("timestamp", body.Timestamp), - zap.Int64("now", now)) - // Send rejection - pkt, _ := protocol.NewPairPacket(protocol.PairReject, 0) - dev.Send(pkt) - return nil - } - // Store timestamp for verification key - p.mu.Lock() - p.pairingTimestamp[dev.ID()] = body.Timestamp - p.mu.Unlock() - } - - p.logger.Info("incoming pair request", zap.String("device_id", dev.ID())) - - var vKey string - peerCert := dev.PeerCert() - if peerCert != nil { - vKey = cert.VerificationKey(p.localCert, peerCert, p.pairingTimestampFor(dev.ID())) - p.logger.Info("pairing verification code", - zap.String("device_id", dev.ID()), - zap.String("code", vKey)) - } - - // Set state and wait for user to accept via CLI - dev.SetState(device.StatePairRequestedByPeer) - if p.onStateChanged != nil { - p.onStateChanged() - } - - p.emit(events.TypePairRequested, dev, vKey) - } - - return nil -} - -func (p *PairPlugin) handleUnpairRequest(_ context.Context, dev *device.Device) error { - state := dev.State() - - switch state { - case device.StatePairRequested: - // We requested, they rejected - p.logger.Info("pair request rejected by peer", zap.String("device_id", dev.ID())) - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - p.emit(events.TypePairRejected, dev, "") - - case device.StatePairRequestedByPeer: - // They requested, then cancelled - p.logger.Info("pair request cancelled by peer", zap.String("device_id", dev.ID())) - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - p.emit(events.TypePairRejected, dev, "") - - case device.StatePaired: - // Unpair request - p.logger.Info("unpair request received", zap.String("device_id", dev.ID())) - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - - case device.StateUnpaired, device.StateUnknown: - // Already unpaired, ignore - p.logger.Debug("ignoring unpair request for unpaired device", zap.String("device_id", dev.ID())) - } - - // Clean up stored timestamp - p.mu.Lock() - delete(p.pairingTimestamp, dev.ID()) - p.mu.Unlock() - - if p.onStateChanged != nil { - p.onStateChanged() - } - - return nil -} - -// AcceptPairing accepts an incoming pair request. -func (p *PairPlugin) AcceptPairing(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairAccept, 0) - if err != nil { - return err - } - - if err := dev.Send(pkt); err != nil { - p.logger.Error("failed to send pair accept", zap.Error(err)) - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - return err - } - - p.pairingDone(dev) - return nil -} - -// RequestPairing initiates a pairing request to a device. -func (p *PairPlugin) RequestPairing(dev *device.Device) error { - if dev.State() == device.StatePaired { - p.logger.Warn("device already paired", zap.String("device_id", dev.ID())) - return nil - } - - if dev.State() == device.StatePairRequestedByPeer { - // They already requested, just accept - return p.AcceptPairing(dev) - } - - // The request timestamp seeds the verification code on both sides, - // so generate it once and send exactly what we store. - timestamp := time.Now().Unix() - pkt, err := protocol.NewPairPacket(protocol.PairAccept, timestamp) - if err != nil { - return err - } - - p.mu.Lock() - p.pairingTimestamp[dev.ID()] = timestamp - p.mu.Unlock() - - if err := dev.Send(pkt); err != nil { - p.logger.Error("failed to send pair request", zap.Error(err)) - return err - } - - peerCert := dev.PeerCert() - if peerCert != nil { - vKey := cert.VerificationKey(p.localCert, peerCert, timestamp) - p.logger.Info("pairing verification code", - zap.String("device_id", dev.ID()), - zap.String("code", vKey)) - } - - dev.SetState(device.StatePairRequested) - p.logger.Info("pair request sent", zap.String("device_id", dev.ID())) - - if p.onStateChanged != nil { - p.onStateChanged() - } - - return nil -} - -// RejectPairing rejects an incoming pair request. -func (p *PairPlugin) RejectPairing(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) - if err != nil { - return err - } - - dev.Send(pkt) // best effort - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - - p.mu.Lock() - delete(p.pairingTimestamp, dev.ID()) - p.mu.Unlock() - - if p.onStateChanged != nil { - p.onStateChanged() - } - - p.logger.Info("pair request rejected", zap.String("device_id", dev.ID())) - return nil -} - -// Unpair removes pairing with a device. -func (p *PairPlugin) Unpair(dev *device.Device) error { - pkt, err := protocol.NewPairPacket(protocol.PairReject, 0) - if err != nil { - return err - } - - dev.Send(pkt) // best effort - dev.SetState(device.StateUnpaired) - dev.ClearEphemeral() - dev.ClearPairDial() - - p.mu.Lock() - delete(p.pairingTimestamp, dev.ID()) - p.mu.Unlock() - - if p.onStateChanged != nil { - p.onStateChanged() - } - - p.logger.Info("device unpaired", zap.String("device_id", dev.ID())) - return nil -} - -func (p *PairPlugin) pairingDone(dev *device.Device) { - dev.SetState(device.StatePaired) - dev.ClearPairDial() - - p.mu.Lock() - delete(p.pairingTimestamp, dev.ID()) - p.mu.Unlock() - - if p.onStateChanged != nil { - p.onStateChanged() - } - - p.logger.Info("pairing complete", zap.String("device_id", dev.ID())) - p.emit(events.TypePairAccepted, dev, "") -} diff --git a/internal/plugins/pair/types.go b/internal/plugins/pair/types.go new file mode 100644 index 0000000..8a81525 --- /dev/null +++ b/internal/plugins/pair/types.go @@ -0,0 +1,83 @@ +package pair + +import ( + "crypto/x509" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +const ( + // AllowedTimestampDiff is the maximum allowed time difference for pairing timestamps (30 min) + AllowedTimestampDiff = 1800 +) + +// PairPlugin handles KDE Connect pairing protocol. +type PairPlugin struct { + devices *device.Registry + localCert *x509.Certificate + onStateChanged func() // callback to persist state + logger *zap.Logger + bus *events.Bus + cfg config.PairingConfig + + mu sync.Mutex + pairingTimestamp map[string]int64 // deviceID -> timestamp from pair request +} + +// NewPairPlugin creates a new pairing plugin. +func NewPairPlugin(devices *device.Registry, localCert *x509.Certificate, cfg config.PairingConfig, onStateChanged func(), bus *events.Bus, logger *zap.Logger) *PairPlugin { + return &PairPlugin{ + devices: devices, + localCert: localCert, + cfg: cfg, + onStateChanged: onStateChanged, + logger: logger.Named("pair"), + bus: bus, + pairingTimestamp: make(map[string]int64), + } +} + +// pairingTimestampFor returns the stored pair-request timestamp for a +// device, or zero if none was recorded (pre-v8 peer or unknown). +func (p *PairPlugin) pairingTimestampFor(deviceID string) int64 { + p.mu.Lock() + defer p.mu.Unlock() + return p.pairingTimestamp[deviceID] +} + +// emit publishes an event to the bus if one is configured. +func (p *PairPlugin) emit(typ events.EventType, dev *device.Device, vKey string) { + if p.bus == nil { + return + } + payload := map[string]interface{}{ + "name": dev.Name(), + "type": dev.Type, + } + if vKey != "" { + payload["verificationKey"] = vKey + } + p.bus.Publish(typ, dev.ID(), payload) +} + +func (p *PairPlugin) Name() string { return "Pair" } + +func (p *PairPlugin) Timeout() time.Duration { return 5 * time.Second } + +func (p *PairPlugin) IncomingTypes() []string { + return []string{protocol.TypePair} +} + +func (p *PairPlugin) OutgoingTypes() []string { + return []string{protocol.TypePair} +} + +func (p *PairPlugin) OnConnect(dev device.Sender) {} + +func (p *PairPlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/sftp/handle.go b/internal/plugins/sftp/handle.go new file mode 100644 index 0000000..67ff5c6 --- /dev/null +++ b/internal/plugins/sftp/handle.go @@ -0,0 +1,65 @@ +package sftp + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *SftpPlugin) Handle(_ context.Context, dev device.Sender, pkt *protocol.Packet) error { + var body SftpBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return err + } + + if body.ErrorMessage != "" { + p.logger.Warn("SFTP server error from device", + zap.String("device_id", dev.ID()), + zap.String("error", body.ErrorMessage), + ) + if p.bus != nil { + p.bus.Publish(events.TypeSftpMount, dev.ID(), map[string]interface{}{ + "error": body.ErrorMessage, + }) + } + return nil + } + + p.mu.Lock() + p.lastBody[dev.ID()] = body + p.mu.Unlock() + + safeURI := fmt.Sprintf("sftp://%s@%s:%s%s", body.User, body.IP, body.Port.String(), body.Path) + p.logger.Info("SFTP server available", zap.String("uri", safeURI)) + + evtPayload := map[string]interface{}{ + "uri": fmt.Sprintf("sftp://%s:%s@%s:%s%s", body.User, body.Password, body.IP, body.Port.String(), body.Path), + "ip": body.IP, + "port": body.Port.String(), + "user": body.User, + "password": body.Password, + "path": body.Path, + } + if len(body.MultiPaths) > 0 { + volumes := make([]map[string]string, 0, len(body.MultiPaths)) + for i, mp := range body.MultiPaths { + name := mp + if i < len(body.PathNames) { + name = body.PathNames[i] + } + volumes = append(volumes, map[string]string{"name": name, "path": mp}) + } + evtPayload["volumes"] = volumes + } + + if p.bus != nil { + p.bus.Publish(events.TypeSftpMount, dev.ID(), evtPayload) + } + + return nil +} diff --git a/internal/plugins/sftp/mount.go b/internal/plugins/sftp/mount.go new file mode 100644 index 0000000..127d612 --- /dev/null +++ b/internal/plugins/sftp/mount.go @@ -0,0 +1,285 @@ +package sftp + +import ( + "context" + "fmt" + "net" + "os" + "os/exec" + "path/filepath" + "regexp" + "strconv" + "strings" + "syscall" + "time" + + "github.com/bethropolis/kcd/internal/device" + "go.uber.org/zap" +) + +// sshUserPattern allows the generated Android SFTP usernames (alphanumerics, +// underscore, dot, hyphen) while rejecting anything starting with '-' — +// sshfs would parse that as an option flag (e.g. -oProxyCommand=...), +// yielding local command execution. +var sshUserPattern = regexp.MustCompile(`^[A-Za-z0-9_][A-Za-z0-9_.-]*$`) + +// sshHostPattern allows IPs (validated separately) and plain hostnames +// (.local, LAN names). Anything else — flags, spaces, shell metachars, +// userinfo (@) — is rejected. +var sshHostPattern = regexp.MustCompile(`^[A-Za-z0-9]([A-Za-z0-9.-]*[A-Za-z0-9])?$`) + +// buildSSHFSArgs validates the phone-provided (or IPC-provided) remote +// parameters and constructs the sshfs argv. Validation (not a "--" +// separator — unsupported by older sshfs 2.x) is what prevents option +// injection: no validated value can begin with '-', so sshfs/fuse option +// parsing can never reinterpret remoteRoot as a flag like -oProxyCommand. +func buildSSHFSArgs(body SftpBody, remotePath, mountPoint string, uid, gid int, keepaliveInterval, keepaliveCount int, extraOpts []string) ([]string, error) { + if !sshUserPattern.MatchString(body.User) || len(body.User) > 64 { + return nil, fmt.Errorf("sftp: refusing suspicious ssh user %q", body.User) + } + if body.IP == "" || strings.HasPrefix(body.IP, "-") { + return nil, fmt.Errorf("sftp: refusing suspicious ssh host %q", body.IP) + } + if net.ParseIP(body.IP) == nil && (!sshHostPattern.MatchString(body.IP) || len(body.IP) > 253) { + return nil, fmt.Errorf("sftp: refusing invalid ssh host %q", body.IP) + } + port, err := strconv.Atoi(strings.TrimSpace(body.Port.String())) + if err != nil || port < 1 || port > 65535 { + return nil, fmt.Errorf("sftp: refusing invalid ssh port %q", body.Port.String()) + } + if remotePath == "" || strings.HasPrefix(remotePath, "-") { + return nil, fmt.Errorf("sftp: refusing suspicious remote path %q", remotePath) + } + remotePath = filepath.Clean(remotePath) + remoteRoot := fmt.Sprintf("%s@%s:%s", body.User, body.IP, remotePath) + + args := []string{ + remoteRoot, + mountPoint, + "-p", strconv.Itoa(port), + "-s", + "-F", "/dev/null", + "-o", "password_stdin", + "-o", "StrictHostKeyChecking=no", + "-o", "UserKnownHostsFile=/dev/null", + "-o", "reconnect", + "-o", "ServerAliveInterval=" + strconv.Itoa(keepaliveInterval), + "-o", "ServerAliveCountMax=" + strconv.Itoa(keepaliveCount), + "-o", "auto_cache", + "-o", "kernel_cache", + "-o", "uid=" + strconv.Itoa(uid), + "-o", "gid=" + strconv.Itoa(gid), + } + + // ExtraSshfsOpts comes from the local operator config, not the phone — + // passed through as-is. + for _, opt := range extraOpts { + args = append(args, "-o", opt) + } + return args, nil +} + +// mountWithBody performs the sshfs mount and returns the local browse path. +// volumePath specifies which storage volume to mount. If empty, the first +// available volume is selected automatically. +func (p *SftpPlugin) mountWithBody(ctx context.Context, deviceID string, body SftpBody, volumePath string) (string, error) { + baseDir := p.cfg.MountDir + if baseDir == "" { + baseDir = os.TempDir() + } + mountPoint := filepath.Join(baseDir, "kcd-sftp-"+deviceID) + if err := os.MkdirAll(mountPoint, 0700); err != nil { + return "", fmt.Errorf("create mount point %s: %w", mountPoint, err) + } + + // Determine the remote path on the Android device. + // The Android SFTP server exposes the real filesystem at "/". + // Listing "/" via sshfs fails because it contains permission-denied + // entries (/proc, /sys). Instead, mount directly to a storage + // volume (e.g. /storage/emulated/0) which is guaranteed browsable. + // If a specific volumePath is provided, use it; otherwise auto-select + // the first available volume. + remotePath := volumePath + if remotePath == "" { + if len(body.MultiPaths) > 0 { + remotePath = body.MultiPaths[0] + } else if body.Path != "" && body.Path != "/" { + remotePath = body.Path + } + } + args, err := buildSSHFSArgs(body, remotePath, mountPoint, os.Getuid(), os.Getgid(), p.cfg.KeepaliveIntervalSecs, p.cfg.KeepaliveCount, p.cfg.ExtraSshfsOpts) + if err != nil { + _ = os.Remove(mountPoint) + return "", err + } + + cmd := exec.CommandContext(ctx, "sshfs", args...) + cmd.Stdin = strings.NewReader(body.Password + "\n") + + if out, err := cmd.CombinedOutput(); err != nil { + _ = os.Remove(mountPoint) + msg := strings.TrimSpace(string(out)) + errMsg := fmt.Sprintf("sshfs failed: %v\n%s", err, msg) + if strings.Contains(msg, "Operation not permitted") || strings.Contains(msg, "fusermount") { + errMsg += "\n\nHint: FUSE requires user_allow_other in /etc/fuse.conf.\nRun: sudo sed -i 's/^#user_allow_other/user_allow_other/' /etc/fuse.conf" + } else if strings.Contains(msg, "sshfs: not found") || strings.Contains(msg, "executable file not found") { + errMsg += "\n\nHint: sshfs is not installed.\nInstall: sudo apt install sshfs (or the equivalent for your distro)" + } + return "", fmt.Errorf("%s", errMsg) + } + + // The mount point now IS the storage volume root, so the user browses + // directly to the mount point — no extra navigation needed. + browsePath := mountPoint + + // Track the mount point so Unmount() can call fusermount. + p.mu.Lock() + p.mountPoints[deviceID] = mountPoint + p.mu.Unlock() + + // Find and track the sshfs daemon PID for graceful shutdown. + if pid, err := findSSHFSPID(mountPoint); err == nil { + p.mu.Lock() + p.mountPIDs[deviceID] = pid + p.mu.Unlock() + p.logger.Debug("tracking sshfs PID", zap.Int("pid", pid)) + } else { + p.logger.Debug("could not find sshfs PID", zap.Error(err)) + } + + p.logger.Info("SFTP mounted", + zap.String("mount_point", mountPoint), + zap.String("browse_path", browsePath), + ) + + // Open in the default file manager (best effort, non-blocking). + if p.cfg.AutoOpen { + go func() { + cmd := p.cfg.OpenCommand + if cmd == "" { + cmd = "xdg-open" + } + if err := exec.CommandContext(context.Background(), cmd, browsePath).Start(); err != nil { + p.logger.Debug("auto-open failed", zap.String("command", cmd), zap.Error(err)) + } + }() + } + + return browsePath, nil +} + +func (p *SftpPlugin) OnConnect(_ device.Sender) {} + +func (p *SftpPlugin) OnDisconnect(dev device.Sender) { + p.mu.Lock() + deviceID := dev.ID() + _, mounted := p.mountPoints[deviceID] + p.mu.Unlock() + + if mounted { + p.logger.Info("device disconnected, cleaning up SFTP mount", + zap.String("device_id", deviceID), + ) + if err := p.Unmount(deviceID); err != nil { + p.logger.Warn("failed to unmount on disconnect", + zap.String("device_id", deviceID), + zap.Error(err), + ) + } + } + + // Evict cached credentials to prevent slow memory leak. + p.mu.Lock() + delete(p.lastBody, deviceID) + p.mu.Unlock() +} + +// Unmount cleanly unmounts a previously mounted SFTP filesystem. +// It first attempts a graceful shutdown of the sshfs process (SIGTERM → wait → SIGKILL), +// then uses fusermount to ensure the mount point is released. +// Returns an error if the device was never mounted. +func (p *SftpPlugin) Unmount(deviceID string) error { + p.mu.Lock() + mountPoint, ok := p.mountPoints[deviceID] + if ok { + delete(p.mountPoints, deviceID) + } + pid, hasPID := p.mountPIDs[deviceID] + if hasPID { + delete(p.mountPIDs, deviceID) + } + p.mu.Unlock() + + if !ok { + return fmt.Errorf("no active SFTP mount for device %s", deviceID) + } + + p.logger.Info("unmounting SFTP share", zap.String("mount_point", mountPoint)) + + // Graceful shutdown: SIGTERM → wait → SIGKILL. + if hasPID { + p.logger.Debug("sending SIGTERM to sshfs", zap.Int("pid", pid)) + proc, err := os.FindProcess(pid) + if err == nil { + if err := proc.Signal(syscall.SIGTERM); err == nil { + done := make(chan struct{}) + go func() { + proc.Wait() + close(done) + }() + select { + case <-done: + p.logger.Debug("sshfs exited cleanly after SIGTERM") + case <-time.After(3 * time.Second): + p.logger.Debug("sshfs did not exit after SIGTERM, sending SIGKILL") + proc.Kill() + } + } + } + } + + // Ensure the mount point is released. + tool := "fusermount3" + if _, err := exec.LookPath(tool); err != nil { + tool = "fusermount" + } + + if out, err := exec.CommandContext(context.Background(), tool, "-u", mountPoint).CombinedOutput(); err != nil { + p.logger.Warn("fusermount cleanup failed", + zap.String("mount_point", mountPoint), + zap.Error(err), + zap.String("output", strings.TrimSpace(string(out))), + ) + } + + _ = os.Remove(mountPoint) + p.logger.Info("SFTP unmounted", zap.String("mount_point", mountPoint)) + return nil +} + +// findSSHFSPID scans /proc to find the sshfs daemon PID for a given mount point. +// Uses /proc directly to avoid external dependencies (pgrep, etc.). +func findSSHFSPID(mountPoint string) (int, error) { + entries, err := os.ReadDir("/proc") + if err != nil { + return 0, fmt.Errorf("read /proc: %w", err) + } + for _, e := range entries { + if !e.IsDir() { + continue + } + pid, err := strconv.Atoi(e.Name()) + if err != nil { + continue + } + cmdline, err := os.ReadFile(filepath.Join("/proc", e.Name(), "cmdline")) + if err != nil { + continue + } + // cmdline uses null bytes as separators; convert to string for matching. + if strings.Contains(string(cmdline), mountPoint) && strings.Contains(string(cmdline), "sshfs") { + return pid, nil + } + } + return 0, fmt.Errorf("no sshfs process found for mount point %s", mountPoint) +} diff --git a/internal/plugins/sftp/request.go b/internal/plugins/sftp/request.go new file mode 100644 index 0000000..70dff91 --- /dev/null +++ b/internal/plugins/sftp/request.go @@ -0,0 +1,205 @@ +package sftp + +import ( + "context" + "fmt" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// RequestMount sends a kdeconnect.sftp.request packet asking the device to +// start its SFTP server and return credentials. +func (p *SftpPlugin) RequestMount(dev device.Sender) error { + pkt, err := protocol.NewPacket("kdeconnect.sftp.request", map[string]any{ + "startBrowsing": true, + }) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// RequestAndMount sends the SFTP request, waits for the Android device to +// respond with credentials (up to 20 s), mounts the filesystem via sshfs, +// and returns the local path the user should open. +func (p *SftpPlugin) RequestAndMount(ctx context.Context, dev device.Sender) (string, error) { + if p.bus == nil { + return "", fmt.Errorf("event bus not available") + } + + // Subscribe BEFORE sending the request to guarantee we don't miss the response. + sub := p.bus.Subscribe(0, events.TypeSftpMount) + defer sub.Close() + + if err := p.RequestMount(dev); err != nil { + return "", fmt.Errorf("send SFTP request: %w", err) + } + + p.logger.Info("SFTP request sent, waiting for phone response", zap.String("device", dev.ID())) + + timeout := time.Duration(p.cfg.CredentialsTimeoutSecs) * time.Second + if timeout == 0 { + timeout = 20 * time.Second + } + deadline, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + for { + select { + case evt, ok := <-sub.C: + if !ok { + return "", fmt.Errorf("event bus closed") + } + if evt.DeviceID != dev.ID() { + continue + } + p.mu.RLock() + body, exists := p.lastBody[dev.ID()] + p.mu.RUnlock() + if !exists { + return "", fmt.Errorf("credentials missing after event (internal error)") + } + return p.mountWithBody(ctx, dev.ID(), body, "") + + case <-deadline.Done(): + return "", fmt.Errorf("timed out after %s waiting for SFTP response — is the KDE Connect app open on the phone?", timeout) + } + } +} + +// RequestAndMountVolume sends the SFTP request, waits for credentials, then +// mounts the specified volume. If volumePath is empty, the available volumes +// are returned without mounting (list mode). The caller is responsible for +// closing the returned closer when done with the mounted path. +func (p *SftpPlugin) RequestAndMountVolume(ctx context.Context, dev device.Sender, volumePath string) (mountPath string, volumes []StorageVolume, err error) { + if p.bus == nil { + return "", nil, fmt.Errorf("event bus not available") + } + + sub := p.bus.Subscribe(0, events.TypeSftpMount) + defer sub.Close() + + if err := p.RequestMount(dev); err != nil { + return "", nil, fmt.Errorf("send SFTP request: %w", err) + } + + p.logger.Info("SFTP request sent, waiting for phone response", zap.String("device", dev.ID())) + + timeout := time.Duration(p.cfg.CredentialsTimeoutSecs) * time.Second + if timeout == 0 { + timeout = 20 * time.Second + } + deadline, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + for { + select { + case evt, ok := <-sub.C: + if !ok { + return "", nil, fmt.Errorf("event bus closed") + } + if evt.DeviceID != dev.ID() { + continue + } + p.mu.RLock() + body, exists := p.lastBody[dev.ID()] + p.mu.RUnlock() + if !exists { + return "", nil, fmt.Errorf("credentials missing after event (internal error)") + } + + vols := p.buildVolumes(body) + + if volumePath == "" { + return "", vols, nil + } + + path, err := p.mountWithBody(ctx, dev.ID(), body, volumePath) + if err != nil { + return "", nil, err + } + return path, vols, nil + + case <-deadline.Done(): + return "", nil, fmt.Errorf("timed out after %s waiting for SFTP response — is the KDE Connect app open on the phone?", timeout) + } + } +} + +// MountLocally mounts using previously cached credentials. +// Prefer RequestAndMount for a one-step experience. +func (p *SftpPlugin) MountLocally(ctx context.Context, deviceID string) (string, error) { + p.mu.RLock() + body, ok := p.lastBody[deviceID] + p.mu.RUnlock() + if !ok { + return "", fmt.Errorf("no SFTP credentials cached for device %s — use 'kcd sftp mount' which requests them automatically", deviceID) + } + return p.mountWithBody(ctx, deviceID, body, "") +} + +// Info returns the cached SFTP connection details for a device. +// Returns nil if no credentials have been received yet. +func (p *SftpPlugin) Info(deviceID string) *SftpInfo { + p.mu.RLock() + defer p.mu.RUnlock() + body, ok := p.lastBody[deviceID] + if !ok { + return nil + } + info := &SftpInfo{ + IP: body.IP, + Port: body.Port, + User: body.User, + Password: body.Password, + Path: body.Path, + } + for i, mp := range body.MultiPaths { + name := mp + if i < len(body.PathNames) { + name = body.PathNames[i] + } + info.Volumes = append(info.Volumes, StorageVolume{Name: name, Path: mp}) + } + return info +} + +// buildVolumes constructs a StorageVolume slice from a SftpBody. +// Caller must hold at least a read lock on p.mu if body comes from p.lastBody. +func (p *SftpPlugin) buildVolumes(body SftpBody) []StorageVolume { + if len(body.MultiPaths) == 0 { + return nil + } + volumes := make([]StorageVolume, 0, len(body.MultiPaths)) + for i, mp := range body.MultiPaths { + name := mp + if i < len(body.PathNames) { + name = body.PathNames[i] + } + volumes = append(volumes, StorageVolume{Name: name, Path: mp}) + } + return volumes +} + +// Volumes returns the list of available storage volumes from cached credentials. +// Returns nil if no credentials or no multiPaths data. +func (p *SftpPlugin) Volumes(deviceID string) []StorageVolume { + p.mu.RLock() + defer p.mu.RUnlock() + body, ok := p.lastBody[deviceID] + if !ok { + return nil + } + return p.buildVolumes(body) +} + +// MountedPath returns the local mount point for a device, or "" if not mounted. +func (p *SftpPlugin) MountedPath(deviceID string) string { + p.mu.RLock() + defer p.mu.RUnlock() + return p.mountPoints[deviceID] +} diff --git a/internal/plugins/sftp/sftp.go b/internal/plugins/sftp/sftp.go deleted file mode 100644 index cee18b7..0000000 --- a/internal/plugins/sftp/sftp.go +++ /dev/null @@ -1,601 +0,0 @@ -package sftp - -import ( - "context" - "encoding/json" - "fmt" - "net" - "os" - "os/exec" - "path/filepath" - "regexp" - "strconv" - "strings" - "sync" - "syscall" - "time" - - "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "go.uber.org/zap" -) - -// SftpPlugin handles KDE Connect SFTP negotiation and optional sshfs mounting. -type SftpPlugin struct { - cfg config.SFTPConfig - bus *events.Bus - logger *zap.Logger - mu sync.RWMutex - lastBody map[string]SftpBody - mountPoints map[string]string // deviceID -> local mountPoint path - mountPIDs map[string]int // deviceID -> sshfs PID for graceful shutdown -} - -func NewSftpPlugin(cfg config.SFTPConfig, bus *events.Bus, logger *zap.Logger) *SftpPlugin { - return &SftpPlugin{ - cfg: cfg, - bus: bus, - logger: logger.With(zap.String("plugin", "sftp")), - lastBody: make(map[string]SftpBody), - mountPoints: make(map[string]string), - mountPIDs: make(map[string]int), - } -} - -// SftpBody matches the body of a kdeconnect.sftp packet sent by the Android app. -type SftpBody struct { - IP string `json:"ip"` - Port json.Number `json:"port"` - User string `json:"user"` - // Password is intentionally not logged. - Password string `json:"password"` - // Path is the primary storage root path from the Android device. - // When exactly one volume exists this is the volume path (e.g. /storage/emulated/0); - // when multiple volumes exist this falls back to "/" (legacy compat). - // Prefer MultiPaths for the authoritative list of browsable roots. - Path string `json:"path"` - // MultiPaths lists all available storage root paths on the device - // (e.g. internal storage, SD card). Populated by Android API 30+. - MultiPaths []string `json:"multiPaths,omitempty"` - // PathNames provides human-readable labels for each path in MultiPaths. - PathNames []string `json:"pathNames,omitempty"` - // ErrorMessage is set when the device cannot start the SFTP server - // (e.g. missing storage permissions). - ErrorMessage string `json:"errorMessage,omitempty"` -} - -// StorageVolume describes a single browsable storage root on the device. -type StorageVolume struct { - Name string `json:"name"` - Path string `json:"path"` -} - -// SftpInfo holds the complete cached SFTP connection details for a device. -type SftpInfo struct { - IP string `json:"ip"` - Port json.Number `json:"port"` - User string `json:"user"` - Password string `json:"password"` - Path string `json:"path"` - Volumes []StorageVolume `json:"volumes,omitempty"` -} - -func (p *SftpPlugin) Name() string { return "SFTP" } -func (p *SftpPlugin) Timeout() time.Duration { return 5 * time.Second } -func (p *SftpPlugin) IncomingTypes() []string { return []string{"kdeconnect.sftp"} } -func (p *SftpPlugin) OutgoingTypes() []string { return []string{"kdeconnect.sftp.request"} } - -func (p *SftpPlugin) Handle(_ context.Context, dev device.Sender, pkt *protocol.Packet) error { - var body SftpBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return err - } - - if body.ErrorMessage != "" { - p.logger.Warn("SFTP server error from device", - zap.String("device_id", dev.ID()), - zap.String("error", body.ErrorMessage), - ) - if p.bus != nil { - p.bus.Publish(events.TypeSftpMount, dev.ID(), map[string]interface{}{ - "error": body.ErrorMessage, - }) - } - return nil - } - - p.mu.Lock() - p.lastBody[dev.ID()] = body - p.mu.Unlock() - - safeURI := fmt.Sprintf("sftp://%s@%s:%s%s", body.User, body.IP, body.Port.String(), body.Path) - p.logger.Info("SFTP server available", zap.String("uri", safeURI)) - - evtPayload := map[string]interface{}{ - "uri": fmt.Sprintf("sftp://%s:%s@%s:%s%s", body.User, body.Password, body.IP, body.Port.String(), body.Path), - "ip": body.IP, - "port": body.Port.String(), - "user": body.User, - "password": body.Password, - "path": body.Path, - } - if len(body.MultiPaths) > 0 { - volumes := make([]map[string]string, 0, len(body.MultiPaths)) - for i, mp := range body.MultiPaths { - name := mp - if i < len(body.PathNames) { - name = body.PathNames[i] - } - volumes = append(volumes, map[string]string{"name": name, "path": mp}) - } - evtPayload["volumes"] = volumes - } - - if p.bus != nil { - p.bus.Publish(events.TypeSftpMount, dev.ID(), evtPayload) - } - - return nil -} - -// RequestMount sends a kdeconnect.sftp.request packet asking the device to -// start its SFTP server and return credentials. -func (p *SftpPlugin) RequestMount(dev device.Sender) error { - pkt, err := protocol.NewPacket("kdeconnect.sftp.request", map[string]any{ - "startBrowsing": true, - }) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// RequestAndMount sends the SFTP request, waits for the Android device to -// respond with credentials (up to 20 s), mounts the filesystem via sshfs, -// and returns the local path the user should open. -func (p *SftpPlugin) RequestAndMount(ctx context.Context, dev device.Sender) (string, error) { - if p.bus == nil { - return "", fmt.Errorf("event bus not available") - } - - // Subscribe BEFORE sending the request to guarantee we don't miss the response. - sub := p.bus.Subscribe(0, events.TypeSftpMount) - defer sub.Close() - - if err := p.RequestMount(dev); err != nil { - return "", fmt.Errorf("send SFTP request: %w", err) - } - - p.logger.Info("SFTP request sent, waiting for phone response", zap.String("device", dev.ID())) - - timeout := time.Duration(p.cfg.CredentialsTimeoutSecs) * time.Second - if timeout == 0 { - timeout = 20 * time.Second - } - deadline, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - - for { - select { - case evt, ok := <-sub.C: - if !ok { - return "", fmt.Errorf("event bus closed") - } - if evt.DeviceID != dev.ID() { - continue - } - p.mu.RLock() - body, exists := p.lastBody[dev.ID()] - p.mu.RUnlock() - if !exists { - return "", fmt.Errorf("credentials missing after event (internal error)") - } - return p.mountWithBody(ctx, dev.ID(), body, "") - - case <-deadline.Done(): - return "", fmt.Errorf("timed out after %s waiting for SFTP response — is the KDE Connect app open on the phone?", timeout) - } - } -} - -// RequestAndMountVolume sends the SFTP request, waits for credentials, then -// mounts the specified volume. If volumePath is empty, the available volumes -// are returned without mounting (list mode). The caller is responsible for -// closing the returned closer when done with the mounted path. -func (p *SftpPlugin) RequestAndMountVolume(ctx context.Context, dev device.Sender, volumePath string) (mountPath string, volumes []StorageVolume, err error) { - if p.bus == nil { - return "", nil, fmt.Errorf("event bus not available") - } - - sub := p.bus.Subscribe(0, events.TypeSftpMount) - defer sub.Close() - - if err := p.RequestMount(dev); err != nil { - return "", nil, fmt.Errorf("send SFTP request: %w", err) - } - - p.logger.Info("SFTP request sent, waiting for phone response", zap.String("device", dev.ID())) - - timeout := time.Duration(p.cfg.CredentialsTimeoutSecs) * time.Second - if timeout == 0 { - timeout = 20 * time.Second - } - deadline, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - - for { - select { - case evt, ok := <-sub.C: - if !ok { - return "", nil, fmt.Errorf("event bus closed") - } - if evt.DeviceID != dev.ID() { - continue - } - p.mu.RLock() - body, exists := p.lastBody[dev.ID()] - p.mu.RUnlock() - if !exists { - return "", nil, fmt.Errorf("credentials missing after event (internal error)") - } - - vols := p.buildVolumes(body) - - if volumePath == "" { - return "", vols, nil - } - - path, err := p.mountWithBody(ctx, dev.ID(), body, volumePath) - if err != nil { - return "", nil, err - } - return path, vols, nil - - case <-deadline.Done(): - return "", nil, fmt.Errorf("timed out after %s waiting for SFTP response — is the KDE Connect app open on the phone?", timeout) - } - } -} - -// MountLocally mounts using previously cached credentials. -// Prefer RequestAndMount for a one-step experience. -func (p *SftpPlugin) MountLocally(ctx context.Context, deviceID string) (string, error) { - p.mu.RLock() - body, ok := p.lastBody[deviceID] - p.mu.RUnlock() - if !ok { - return "", fmt.Errorf("no SFTP credentials cached for device %s — use 'kcd sftp mount' which requests them automatically", deviceID) - } - return p.mountWithBody(ctx, deviceID, body, "") -} - -// sshUserPattern allows the generated Android SFTP usernames (alphanumerics, -// underscore, dot, hyphen) while rejecting anything starting with '-' — -// sshfs would parse that as an option flag (e.g. -oProxyCommand=...), -// yielding local command execution. -var sshUserPattern = regexp.MustCompile(`^[A-Za-z0-9_][A-Za-z0-9_.-]*$`) - -// sshHostPattern allows IPs (validated separately) and plain hostnames -// (.local, LAN names). Anything else — flags, spaces, shell metachars, -// userinfo (@) — is rejected. -var sshHostPattern = regexp.MustCompile(`^[A-Za-z0-9]([A-Za-z0-9.-]*[A-Za-z0-9])?$`) - -// buildSSHFSArgs validates the phone-provided (or IPC-provided) remote -// parameters and constructs the sshfs argv. Validation (not a "--" -// separator — unsupported by older sshfs 2.x) is what prevents option -// injection: no validated value can begin with '-', so sshfs/fuse option -// parsing can never reinterpret remoteRoot as a flag like -oProxyCommand. -func buildSSHFSArgs(body SftpBody, remotePath, mountPoint string, uid, gid int, keepaliveInterval, keepaliveCount int, extraOpts []string) ([]string, error) { - if !sshUserPattern.MatchString(body.User) || len(body.User) > 64 { - return nil, fmt.Errorf("sftp: refusing suspicious ssh user %q", body.User) - } - if body.IP == "" || strings.HasPrefix(body.IP, "-") { - return nil, fmt.Errorf("sftp: refusing suspicious ssh host %q", body.IP) - } - if net.ParseIP(body.IP) == nil && (!sshHostPattern.MatchString(body.IP) || len(body.IP) > 253) { - return nil, fmt.Errorf("sftp: refusing invalid ssh host %q", body.IP) - } - port, err := strconv.Atoi(strings.TrimSpace(body.Port.String())) - if err != nil || port < 1 || port > 65535 { - return nil, fmt.Errorf("sftp: refusing invalid ssh port %q", body.Port.String()) - } - if remotePath == "" || strings.HasPrefix(remotePath, "-") { - return nil, fmt.Errorf("sftp: refusing suspicious remote path %q", remotePath) - } - remotePath = filepath.Clean(remotePath) - remoteRoot := fmt.Sprintf("%s@%s:%s", body.User, body.IP, remotePath) - - args := []string{ - remoteRoot, - mountPoint, - "-p", strconv.Itoa(port), - "-s", - "-F", "/dev/null", - "-o", "password_stdin", - "-o", "StrictHostKeyChecking=no", - "-o", "UserKnownHostsFile=/dev/null", - "-o", "reconnect", - "-o", "ServerAliveInterval=" + strconv.Itoa(keepaliveInterval), - "-o", "ServerAliveCountMax=" + strconv.Itoa(keepaliveCount), - "-o", "auto_cache", - "-o", "kernel_cache", - "-o", "uid=" + strconv.Itoa(uid), - "-o", "gid=" + strconv.Itoa(gid), - } - - // ExtraSshfsOpts comes from the local operator config, not the phone — - // passed through as-is. - for _, opt := range extraOpts { - args = append(args, "-o", opt) - } - return args, nil -} - -// mountWithBody performs the sshfs mount and returns the local browse path. -// volumePath specifies which storage volume to mount. If empty, the first -// available volume is selected automatically. -func (p *SftpPlugin) mountWithBody(ctx context.Context, deviceID string, body SftpBody, volumePath string) (string, error) { - baseDir := p.cfg.MountDir - if baseDir == "" { - baseDir = os.TempDir() - } - mountPoint := filepath.Join(baseDir, "kcd-sftp-"+deviceID) - if err := os.MkdirAll(mountPoint, 0700); err != nil { - return "", fmt.Errorf("create mount point %s: %w", mountPoint, err) - } - - // Determine the remote path on the Android device. - // The Android SFTP server exposes the real filesystem at "/". - // Listing "/" via sshfs fails because it contains permission-denied - // entries (/proc, /sys). Instead, mount directly to a storage - // volume (e.g. /storage/emulated/0) which is guaranteed browsable. - // If a specific volumePath is provided, use it; otherwise auto-select - // the first available volume. - remotePath := volumePath - if remotePath == "" { - if len(body.MultiPaths) > 0 { - remotePath = body.MultiPaths[0] - } else if body.Path != "" && body.Path != "/" { - remotePath = body.Path - } - } - args, err := buildSSHFSArgs(body, remotePath, mountPoint, os.Getuid(), os.Getgid(), p.cfg.KeepaliveIntervalSecs, p.cfg.KeepaliveCount, p.cfg.ExtraSshfsOpts) - if err != nil { - _ = os.Remove(mountPoint) - return "", err - } - - cmd := exec.CommandContext(ctx, "sshfs", args...) - cmd.Stdin = strings.NewReader(body.Password + "\n") - - if out, err := cmd.CombinedOutput(); err != nil { - _ = os.Remove(mountPoint) - msg := strings.TrimSpace(string(out)) - errMsg := fmt.Sprintf("sshfs failed: %v\n%s", err, msg) - if strings.Contains(msg, "Operation not permitted") || strings.Contains(msg, "fusermount") { - errMsg += "\n\nHint: FUSE requires user_allow_other in /etc/fuse.conf.\nRun: sudo sed -i 's/^#user_allow_other/user_allow_other/' /etc/fuse.conf" - } else if strings.Contains(msg, "sshfs: not found") || strings.Contains(msg, "executable file not found") { - errMsg += "\n\nHint: sshfs is not installed.\nInstall: sudo apt install sshfs (or the equivalent for your distro)" - } - return "", fmt.Errorf("%s", errMsg) - } - - // The mount point now IS the storage volume root, so the user browses - // directly to the mount point — no extra navigation needed. - browsePath := mountPoint - - // Track the mount point so Unmount() can call fusermount. - p.mu.Lock() - p.mountPoints[deviceID] = mountPoint - p.mu.Unlock() - - // Find and track the sshfs daemon PID for graceful shutdown. - if pid, err := findSSHFSPID(mountPoint); err == nil { - p.mu.Lock() - p.mountPIDs[deviceID] = pid - p.mu.Unlock() - p.logger.Debug("tracking sshfs PID", zap.Int("pid", pid)) - } else { - p.logger.Debug("could not find sshfs PID", zap.Error(err)) - } - - p.logger.Info("SFTP mounted", - zap.String("mount_point", mountPoint), - zap.String("browse_path", browsePath), - ) - - // Open in the default file manager (best effort, non-blocking). - if p.cfg.AutoOpen { - go func() { - cmd := p.cfg.OpenCommand - if cmd == "" { - cmd = "xdg-open" - } - if err := exec.CommandContext(context.Background(), cmd, browsePath).Start(); err != nil { - p.logger.Debug("auto-open failed", zap.String("command", cmd), zap.Error(err)) - } - }() - } - - return browsePath, nil -} - -// Info returns the cached SFTP connection details for a device. -// Returns nil if no credentials have been received yet. -func (p *SftpPlugin) Info(deviceID string) *SftpInfo { - p.mu.RLock() - defer p.mu.RUnlock() - body, ok := p.lastBody[deviceID] - if !ok { - return nil - } - info := &SftpInfo{ - IP: body.IP, - Port: body.Port, - User: body.User, - Password: body.Password, - Path: body.Path, - } - for i, mp := range body.MultiPaths { - name := mp - if i < len(body.PathNames) { - name = body.PathNames[i] - } - info.Volumes = append(info.Volumes, StorageVolume{Name: name, Path: mp}) - } - return info -} - -// buildVolumes constructs a StorageVolume slice from a SftpBody. -// Caller must hold at least a read lock on p.mu if body comes from p.lastBody. -func (p *SftpPlugin) buildVolumes(body SftpBody) []StorageVolume { - if len(body.MultiPaths) == 0 { - return nil - } - volumes := make([]StorageVolume, 0, len(body.MultiPaths)) - for i, mp := range body.MultiPaths { - name := mp - if i < len(body.PathNames) { - name = body.PathNames[i] - } - volumes = append(volumes, StorageVolume{Name: name, Path: mp}) - } - return volumes -} - -// Volumes returns the list of available storage volumes from cached credentials. -// Returns nil if no credentials or no multiPaths data. -func (p *SftpPlugin) Volumes(deviceID string) []StorageVolume { - p.mu.RLock() - defer p.mu.RUnlock() - body, ok := p.lastBody[deviceID] - if !ok { - return nil - } - return p.buildVolumes(body) -} - -func (p *SftpPlugin) OnConnect(_ device.Sender) {} - -func (p *SftpPlugin) OnDisconnect(dev device.Sender) { - p.mu.Lock() - deviceID := dev.ID() - _, mounted := p.mountPoints[deviceID] - p.mu.Unlock() - - if mounted { - p.logger.Info("device disconnected, cleaning up SFTP mount", - zap.String("device_id", deviceID), - ) - if err := p.Unmount(deviceID); err != nil { - p.logger.Warn("failed to unmount on disconnect", - zap.String("device_id", deviceID), - zap.Error(err), - ) - } - } - - // Evict cached credentials to prevent slow memory leak. - p.mu.Lock() - delete(p.lastBody, deviceID) - p.mu.Unlock() -} - -// Unmount cleanly unmounts a previously mounted SFTP filesystem. -// It first attempts a graceful shutdown of the sshfs process (SIGTERM → wait → SIGKILL), -// then uses fusermount to ensure the mount point is released. -// Returns an error if the device was never mounted. -func (p *SftpPlugin) Unmount(deviceID string) error { - p.mu.Lock() - mountPoint, ok := p.mountPoints[deviceID] - if ok { - delete(p.mountPoints, deviceID) - } - pid, hasPID := p.mountPIDs[deviceID] - if hasPID { - delete(p.mountPIDs, deviceID) - } - p.mu.Unlock() - - if !ok { - return fmt.Errorf("no active SFTP mount for device %s", deviceID) - } - - p.logger.Info("unmounting SFTP share", zap.String("mount_point", mountPoint)) - - // Graceful shutdown: SIGTERM → wait → SIGKILL. - if hasPID { - p.logger.Debug("sending SIGTERM to sshfs", zap.Int("pid", pid)) - proc, err := os.FindProcess(pid) - if err == nil { - if err := proc.Signal(syscall.SIGTERM); err == nil { - done := make(chan struct{}) - go func() { - proc.Wait() - close(done) - }() - select { - case <-done: - p.logger.Debug("sshfs exited cleanly after SIGTERM") - case <-time.After(3 * time.Second): - p.logger.Debug("sshfs did not exit after SIGTERM, sending SIGKILL") - proc.Kill() - } - } - } - } - - // Ensure the mount point is released. - tool := "fusermount3" - if _, err := exec.LookPath(tool); err != nil { - tool = "fusermount" - } - - if out, err := exec.CommandContext(context.Background(), tool, "-u", mountPoint).CombinedOutput(); err != nil { - p.logger.Warn("fusermount cleanup failed", - zap.String("mount_point", mountPoint), - zap.Error(err), - zap.String("output", strings.TrimSpace(string(out))), - ) - } - - _ = os.Remove(mountPoint) - p.logger.Info("SFTP unmounted", zap.String("mount_point", mountPoint)) - return nil -} - -// MountedPath returns the local mount point for a device, or "" if not mounted. -func (p *SftpPlugin) MountedPath(deviceID string) string { - p.mu.RLock() - defer p.mu.RUnlock() - return p.mountPoints[deviceID] -} - -// findSSHFSPID scans /proc to find the sshfs daemon PID for a given mount point. -// Uses /proc directly to avoid external dependencies (pgrep, etc.). -func findSSHFSPID(mountPoint string) (int, error) { - entries, err := os.ReadDir("/proc") - if err != nil { - return 0, fmt.Errorf("read /proc: %w", err) - } - for _, e := range entries { - if !e.IsDir() { - continue - } - pid, err := strconv.Atoi(e.Name()) - if err != nil { - continue - } - cmdline, err := os.ReadFile(filepath.Join("/proc", e.Name(), "cmdline")) - if err != nil { - continue - } - // cmdline uses null bytes as separators; convert to string for matching. - if strings.Contains(string(cmdline), mountPoint) && strings.Contains(string(cmdline), "sshfs") { - return pid, nil - } - } - return 0, fmt.Errorf("no sshfs process found for mount point %s", mountPoint) -} diff --git a/internal/plugins/sftp/types.go b/internal/plugins/sftp/types.go new file mode 100644 index 0000000..0d49f94 --- /dev/null +++ b/internal/plugins/sftp/types.go @@ -0,0 +1,76 @@ +package sftp + +import ( + "encoding/json" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/events" + "go.uber.org/zap" +) + +// SftpPlugin handles KDE Connect SFTP negotiation and optional sshfs mounting. +type SftpPlugin struct { + cfg config.SFTPConfig + bus *events.Bus + logger *zap.Logger + mu sync.RWMutex + lastBody map[string]SftpBody + mountPoints map[string]string // deviceID -> local mountPoint path + mountPIDs map[string]int // deviceID -> sshfs PID for graceful shutdown +} + +func NewSftpPlugin(cfg config.SFTPConfig, bus *events.Bus, logger *zap.Logger) *SftpPlugin { + return &SftpPlugin{ + cfg: cfg, + bus: bus, + logger: logger.With(zap.String("plugin", "sftp")), + lastBody: make(map[string]SftpBody), + mountPoints: make(map[string]string), + mountPIDs: make(map[string]int), + } +} + +// SftpBody matches the body of a kdeconnect.sftp packet sent by the Android app. +type SftpBody struct { + IP string `json:"ip"` + Port json.Number `json:"port"` + User string `json:"user"` + // Password is intentionally not logged. + Password string `json:"password"` + // Path is the primary storage root path from the Android device. + // When exactly one volume exists this is the volume path (e.g. /storage/emulated/0); + // when multiple volumes exist this falls back to "/" (legacy compat). + // Prefer MultiPaths for the authoritative list of browsable roots. + Path string `json:"path"` + // MultiPaths lists all available storage root paths on the device + // (e.g. internal storage, SD card). Populated by Android API 30+. + MultiPaths []string `json:"multiPaths,omitempty"` + // PathNames provides human-readable labels for each path in MultiPaths. + PathNames []string `json:"pathNames,omitempty"` + // ErrorMessage is set when the device cannot start the SFTP server + // (e.g. missing storage permissions). + ErrorMessage string `json:"errorMessage,omitempty"` +} + +// StorageVolume describes a single browsable storage root on the device. +type StorageVolume struct { + Name string `json:"name"` + Path string `json:"path"` +} + +// SftpInfo holds the complete cached SFTP connection details for a device. +type SftpInfo struct { + IP string `json:"ip"` + Port json.Number `json:"port"` + User string `json:"user"` + Password string `json:"password"` + Path string `json:"path"` + Volumes []StorageVolume `json:"volumes,omitempty"` +} + +func (p *SftpPlugin) Name() string { return "SFTP" } +func (p *SftpPlugin) Timeout() time.Duration { return 5 * time.Second } +func (p *SftpPlugin) IncomingTypes() []string { return []string{"kdeconnect.sftp"} } +func (p *SftpPlugin) OutgoingTypes() []string { return []string{"kdeconnect.sftp.request"} } diff --git a/internal/plugins/share/handle.go b/internal/plugins/share/handle.go new file mode 100644 index 0000000..15a44ec --- /dev/null +++ b/internal/plugins/share/handle.go @@ -0,0 +1,147 @@ +package share + +import ( + "context" + "encoding/json" + "fmt" + "os" + "os/exec" + "path/filepath" + "runtime/debug" + "strings" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *SharePlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + var body ShareBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return fmt.Errorf("share: parse body: %w", err) + } + + if body.Text != "" && pkt.PayloadSize <= 0 { + p.Logger.Info("share: received text", zap.String("text", body.Text)) + if p.bus != nil { + p.bus.Publish(events.TypeShareText, dev.ID(), map[string]string{"text": body.Text}) + } + go func() { + var cmd *exec.Cmd + if os.Getenv("WAYLAND_DISPLAY") != "" { + cmd = exec.CommandContext(context.Background(), "wl-copy") + } else { + cmd = exec.CommandContext(context.Background(), "xclip", "-selection", "clipboard") + } + cmd.Stdin = strings.NewReader(body.Text) + _ = cmd.Run() + }() + return nil + } + + if body.Url != "" && pkt.PayloadSize <= 0 { + p.Logger.Info("share: received url", zap.String("url", body.Url)) + if p.bus != nil { + p.bus.Publish(events.TypeShareURL, dev.ID(), map[string]string{"url": body.Url}) + } + // xdg-open dispatches on URI scheme to arbitrary desktop handlers, + // so only http(s) may reach it: file://, smb:, mailto: and custom + // app schemes would hand phone-influenced input to unrelated local + // handlers. The URL event above still reaches clients either way. + if isOpenableURL(body.Url) { + plugin.RunCommandAsync(p.Logger, "xdg-open", body.Url) + } else { + p.Logger.Warn("share: refusing to open non-http(s) URL", + zap.String("url", body.Url)) + } + return nil + } + + if pkt.PayloadSize <= 0 || pkt.PayloadTransferInfo == nil { + return nil + } + + safeName := SanitizeFilename(body.Filename) + if err := os.MkdirAll(p.DownloadDir, 0755); err != nil { + return fmt.Errorf("share: critical - failed to create download dir %s: %w", p.DownloadDir, err) + } + + destPath, err := EnsureUnique(p.DownloadDir, safeName) + if p.cfg.Overwrite { + destPath = filepath.Join(p.DownloadDir, safeName) + err = nil + } + if err != nil { + return fmt.Errorf("share: collision handling: %w", err) + } + + remoteIP := dev.RemoteIP() + if remoteIP == nil { + return fmt.Errorf("share: failed to resolve remote peer IP") + } + + payloadSize := pkt.PayloadSize + payloadPort := pkt.PayloadTransferInfo.Port + expectedFP := cert.PinnedFingerprint(dev.PeerCert()) + + go func() { + defer debug.FreeOSMemory() + + var onProgress func(int64, int64) + if p.bus != nil { + throttle := newProgressThrottle(p.bus, dev.ID(), body.Filename, payloadSize) + onProgress = throttle.Update + } + + err := ReceiveSideChannel(context.Background(), remoteIP, payloadPort, payloadSize, destPath, p.TLSConfig, expectedFP, onProgress, p.Logger, p.sidechannel) + if err != nil { + p.Logger.Error("share receive failed", zap.Error(err)) + if p.bus != nil { + p.bus.Publish(events.TypeShareComplete, dev.ID(), map[string]interface{}{ + "file": body.Filename, + "success": false, + "error": err.Error(), + }) + } + } else { + if body.LastModified > 0 { + modTime := time.UnixMilli(body.LastModified) + if err := os.Chtimes(destPath, modTime, modTime); err != nil { + p.Logger.Debug("share: failed to restore file timestamps", zap.Error(err)) + } + } + + if p.bus != nil { + p.bus.Publish(events.TypeShareComplete, dev.ID(), map[string]interface{}{ + "file": body.Filename, + "success": true, + }) + } + if p.cfg.AutoOpen { + // Never auto-open executable content: handing .desktop files + // (or scripts) to the desktop handler can execute code. The + // file itself is still saved and announced — open it manually. + if autoOpenBlocked(destPath) { + p.Logger.Warn("share: refusing to auto-open executable file", + zap.String("file", destPath)) + } else { + cmd := p.cfg.OpenCommand + if cmd == "" { + cmd = "xdg-open" + } + absPath, err := filepath.Abs(destPath) + if err != nil { + absPath = destPath + } + plugin.RunCommandAsync(p.Logger, cmd, absPath) + } + } + } + }() + + return nil +} diff --git a/internal/plugins/share/send.go b/internal/plugins/share/send.go new file mode 100644 index 0000000..cb2b09a --- /dev/null +++ b/internal/plugins/share/send.go @@ -0,0 +1,113 @@ +package share + +import ( + "context" + "fmt" + "os" + "path/filepath" + "runtime/debug" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *SharePlugin) SendFile(ctx context.Context, dev device.Sender, filePath string) error { + f, err := os.Open(filePath) + if err != nil { + return fmt.Errorf("share: open file: %w", err) + } + stat, err := f.Stat() + f.Close() // Close immediately; AcceptAndSend opens it again exactly when the phone connects. + if err != nil { + return fmt.Errorf("share: stat file: %w", err) + } + + if stat.IsDir() { + return fmt.Errorf("share: directory transfer is not supported") + } + + // Bind to an available side-channel port (using config range) + ln, port, err := ListenSideChannel(ctx, p.cfg, p.TLSConfig) + if err != nil { + return err + } + + var onProgress func(int64, int64) + if p.bus != nil { + throttle := newProgressThrottle(p.bus, dev.ID(), filepath.Base(filePath), stat.Size()) + onProgress = throttle.Update + } + expectedFP := cert.PinnedFingerprint(dev.PeerCert()) + + // Handle the transfer in the background so IPC returns instantly + go func() { + defer debug.FreeOSMemory() + + timeout := time.Duration(p.cfg.AcceptTimeoutSecs) * time.Second + if timeout == 0 { + timeout = 2 * time.Minute + } + err := AcceptAndSend(ln, filePath, p.TLSConfig, dev.ID(), expectedFP, timeout, onProgress, p.Logger, p.sidechannel) + + if err != nil { + p.Logger.Error("share: send failed", + zap.String("device_id", dev.ID()), + zap.String("file", filepath.Base(filePath)), + zap.Int("port", port), + zap.Error(err), + ) + } else { + p.Logger.Info("share: send complete", + zap.String("device_id", dev.ID()), + zap.String("file", filepath.Base(filePath)), + ) + } + + if p.bus != nil { + payload := map[string]interface{}{ + "file": filepath.Base(filePath), + "success": err == nil, + } + if err != nil { + payload["error"] = err.Error() + } + p.bus.Publish(events.TypeShareComplete, dev.ID(), payload) + } + }() + + modTime := stat.ModTime().UnixMilli() + + // Send invite packet with strict metadata + pkt, err := protocol.NewPacket("kdeconnect.share.request", ShareBody{ + Filename: filepath.Base(filePath), + NumberOfFiles: 1, + TotalPayloadSize: stat.Size(), + LastModified: modTime, + CreationTime: modTime, + }) + if err != nil { + ln.Close() + return err + } + + pkt.PayloadSize = stat.Size() + pkt.PayloadTransferInfo = &protocol.TransferInfo{ + Port: port, + } + + p.Logger.Info("share: sending transfer invitation", + zap.String("device_id", dev.ID()), + zap.String("path", filePath), + zap.Int64("size", pkt.PayloadSize), + zap.Int("port", port), + ) + + return dev.Send(pkt) +} + +func (p *SharePlugin) OnConnect(dev device.Sender) {} +func (p *SharePlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/share/share.go b/internal/plugins/share/share.go deleted file mode 100644 index 3c0da48..0000000 --- a/internal/plugins/share/share.go +++ /dev/null @@ -1,338 +0,0 @@ -package share - -import ( - "context" - "crypto/tls" - "encoding/json" - "fmt" - "os" - "os/exec" - "path/filepath" - "runtime/debug" - "strings" - "sync" - "time" - - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/plugin" - "github.com/bethropolis/kcd/internal/protocol" - "github.com/bethropolis/kcd/internal/transport" - "go.uber.org/zap" -) - -type progressThrottle struct { - bus *events.Bus - deviceID string - filename string - total int64 - interval time.Duration - mu sync.Mutex - last time.Time - pending int64 -} - -func newProgressThrottle(bus *events.Bus, deviceID, filename string, total int64) *progressThrottle { - return &progressThrottle{ - bus: bus, - deviceID: deviceID, - filename: filename, - total: total, - interval: 500 * time.Millisecond, - } -} - -func (t *progressThrottle) Update(current, _ int64) { - if t.bus == nil { - return - } - t.mu.Lock() - t.pending = current - now := time.Now() - if now.Sub(t.last) < t.interval { - t.mu.Unlock() - return - } - t.last = now - cur := t.pending - t.mu.Unlock() - - t.bus.Publish(events.TypeShareProgress, t.deviceID, map[string]any{ - "file": t.filename, - "current": cur, - "total": t.total, - }) -} - -type SharePlugin struct { - sidechannel transport.SidechannelOptions - DownloadDir string - cfg config.ShareConfig - TLSConfig *tls.Config - Logger *zap.Logger - bus *events.Bus -} - -func NewSharePlugin(downloadDir string, cfg config.ShareConfig, tlsConfig *tls.Config, bus *events.Bus, logger *zap.Logger, options ...transport.SidechannelOptions) *SharePlugin { - var sidechannel transport.SidechannelOptions - if len(options) > 0 { - sidechannel = options[0] - } - return &SharePlugin{ - sidechannel: sidechannel, - DownloadDir: downloadDir, - cfg: cfg, - TLSConfig: tlsConfig, - Logger: logger.With(zap.String("plugin", "share")), - bus: bus, - } -} - -// ShareBody includes the missing Android metadata (LastModified/CreationTime) -type ShareBody struct { - Filename string `json:"filename"` - NumberOfFiles int `json:"numberOfFiles,omitempty"` - TotalPayloadSize int64 `json:"totalPayloadSize,omitempty"` - LastModified int64 `json:"lastModified,omitempty"` - CreationTime int64 `json:"creationTime,omitempty"` - Text string `json:"text,omitempty"` - Url string `json:"url,omitempty"` -} - -func (p *SharePlugin) Name() string { return "Share" } - -func (p *SharePlugin) Timeout() time.Duration { return 0 } - -func (p *SharePlugin) IncomingTypes() []string { - return []string{"kdeconnect.share.request"} -} - -func (p *SharePlugin) OutgoingTypes() []string { - return []string{"kdeconnect.share.request"} -} - -func (p *SharePlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - var body ShareBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return fmt.Errorf("share: parse body: %w", err) - } - - if body.Text != "" && pkt.PayloadSize <= 0 { - p.Logger.Info("share: received text", zap.String("text", body.Text)) - if p.bus != nil { - p.bus.Publish(events.TypeShareText, dev.ID(), map[string]string{"text": body.Text}) - } - go func() { - var cmd *exec.Cmd - if os.Getenv("WAYLAND_DISPLAY") != "" { - cmd = exec.CommandContext(context.Background(), "wl-copy") - } else { - cmd = exec.CommandContext(context.Background(), "xclip", "-selection", "clipboard") - } - cmd.Stdin = strings.NewReader(body.Text) - _ = cmd.Run() - }() - return nil - } - - if body.Url != "" && pkt.PayloadSize <= 0 { - p.Logger.Info("share: received url", zap.String("url", body.Url)) - if p.bus != nil { - p.bus.Publish(events.TypeShareURL, dev.ID(), map[string]string{"url": body.Url}) - } - // xdg-open dispatches on URI scheme to arbitrary desktop handlers, - // so only http(s) may reach it: file://, smb:, mailto: and custom - // app schemes would hand phone-influenced input to unrelated local - // handlers. The URL event above still reaches clients either way. - if isOpenableURL(body.Url) { - plugin.RunCommandAsync(p.Logger, "xdg-open", body.Url) - } else { - p.Logger.Warn("share: refusing to open non-http(s) URL", - zap.String("url", body.Url)) - } - return nil - } - - if pkt.PayloadSize <= 0 || pkt.PayloadTransferInfo == nil { - return nil - } - - safeName := SanitizeFilename(body.Filename) - if err := os.MkdirAll(p.DownloadDir, 0755); err != nil { - return fmt.Errorf("share: critical - failed to create download dir %s: %w", p.DownloadDir, err) - } - - destPath, err := EnsureUnique(p.DownloadDir, safeName) - if p.cfg.Overwrite { - destPath = filepath.Join(p.DownloadDir, safeName) - err = nil - } - if err != nil { - return fmt.Errorf("share: collision handling: %w", err) - } - - remoteIP := dev.RemoteIP() - if remoteIP == nil { - return fmt.Errorf("share: failed to resolve remote peer IP") - } - - payloadSize := pkt.PayloadSize - payloadPort := pkt.PayloadTransferInfo.Port - expectedFP := cert.PinnedFingerprint(dev.PeerCert()) - - go func() { - defer debug.FreeOSMemory() - - var onProgress func(int64, int64) - if p.bus != nil { - throttle := newProgressThrottle(p.bus, dev.ID(), body.Filename, payloadSize) - onProgress = throttle.Update - } - - err := ReceiveSideChannel(context.Background(), remoteIP, payloadPort, payloadSize, destPath, p.TLSConfig, expectedFP, onProgress, p.Logger, p.sidechannel) - if err != nil { - p.Logger.Error("share receive failed", zap.Error(err)) - if p.bus != nil { - p.bus.Publish(events.TypeShareComplete, dev.ID(), map[string]interface{}{ - "file": body.Filename, - "success": false, - "error": err.Error(), - }) - } - } else { - if body.LastModified > 0 { - modTime := time.UnixMilli(body.LastModified) - if err := os.Chtimes(destPath, modTime, modTime); err != nil { - p.Logger.Debug("share: failed to restore file timestamps", zap.Error(err)) - } - } - - if p.bus != nil { - p.bus.Publish(events.TypeShareComplete, dev.ID(), map[string]interface{}{ - "file": body.Filename, - "success": true, - }) - } - if p.cfg.AutoOpen { - // Never auto-open executable content: handing .desktop files - // (or scripts) to the desktop handler can execute code. The - // file itself is still saved and announced — open it manually. - if autoOpenBlocked(destPath) { - p.Logger.Warn("share: refusing to auto-open executable file", - zap.String("file", destPath)) - } else { - cmd := p.cfg.OpenCommand - if cmd == "" { - cmd = "xdg-open" - } - absPath, err := filepath.Abs(destPath) - if err != nil { - absPath = destPath - } - plugin.RunCommandAsync(p.Logger, cmd, absPath) - } - } - } - }() - - return nil -} - -func (p *SharePlugin) SendFile(ctx context.Context, dev device.Sender, filePath string) error { - f, err := os.Open(filePath) - if err != nil { - return fmt.Errorf("share: open file: %w", err) - } - stat, err := f.Stat() - f.Close() // Close immediately; AcceptAndSend opens it again exactly when the phone connects. - if err != nil { - return fmt.Errorf("share: stat file: %w", err) - } - - if stat.IsDir() { - return fmt.Errorf("share: directory transfer is not supported") - } - - // Bind to an available side-channel port (using config range) - ln, port, err := ListenSideChannel(ctx, p.cfg, p.TLSConfig) - if err != nil { - return err - } - - var onProgress func(int64, int64) - if p.bus != nil { - throttle := newProgressThrottle(p.bus, dev.ID(), filepath.Base(filePath), stat.Size()) - onProgress = throttle.Update - } - expectedFP := cert.PinnedFingerprint(dev.PeerCert()) - - // Handle the transfer in the background so IPC returns instantly - go func() { - defer debug.FreeOSMemory() - - timeout := time.Duration(p.cfg.AcceptTimeoutSecs) * time.Second - if timeout == 0 { - timeout = 2 * time.Minute - } - err := AcceptAndSend(ln, filePath, p.TLSConfig, dev.ID(), expectedFP, timeout, onProgress, p.Logger, p.sidechannel) - - if err != nil { - p.Logger.Error("share: send failed", - zap.String("device_id", dev.ID()), - zap.String("file", filepath.Base(filePath)), - zap.Int("port", port), - zap.Error(err), - ) - } else { - p.Logger.Info("share: send complete", - zap.String("device_id", dev.ID()), - zap.String("file", filepath.Base(filePath)), - ) - } - - if p.bus != nil { - payload := map[string]interface{}{ - "file": filepath.Base(filePath), - "success": err == nil, - } - if err != nil { - payload["error"] = err.Error() - } - p.bus.Publish(events.TypeShareComplete, dev.ID(), payload) - } - }() - - modTime := stat.ModTime().UnixMilli() - - // Send invite packet with strict metadata - pkt, err := protocol.NewPacket("kdeconnect.share.request", ShareBody{ - Filename: filepath.Base(filePath), - NumberOfFiles: 1, - TotalPayloadSize: stat.Size(), - LastModified: modTime, - CreationTime: modTime, - }) - if err != nil { - ln.Close() - return err - } - - pkt.PayloadSize = stat.Size() - pkt.PayloadTransferInfo = &protocol.TransferInfo{ - Port: port, - } - - p.Logger.Info("share: sending transfer invitation", - zap.String("device_id", dev.ID()), - zap.String("path", filePath), - zap.Int64("size", pkt.PayloadSize), - zap.Int("port", port), - ) - - return dev.Send(pkt) -} - -func (p *SharePlugin) OnConnect(dev device.Sender) {} -func (p *SharePlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/share/types.go b/internal/plugins/share/types.go new file mode 100644 index 0000000..5bd58d6 --- /dev/null +++ b/internal/plugins/share/types.go @@ -0,0 +1,102 @@ +package share + +import ( + "crypto/tls" + "sync" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +type progressThrottle struct { + bus *events.Bus + deviceID string + filename string + total int64 + interval time.Duration + mu sync.Mutex + last time.Time + pending int64 +} + +func newProgressThrottle(bus *events.Bus, deviceID, filename string, total int64) *progressThrottle { + return &progressThrottle{ + bus: bus, + deviceID: deviceID, + filename: filename, + total: total, + interval: 500 * time.Millisecond, + } +} + +func (t *progressThrottle) Update(current, _ int64) { + if t.bus == nil { + return + } + t.mu.Lock() + t.pending = current + now := time.Now() + if now.Sub(t.last) < t.interval { + t.mu.Unlock() + return + } + t.last = now + cur := t.pending + t.mu.Unlock() + + t.bus.Publish(events.TypeShareProgress, t.deviceID, map[string]any{ + "file": t.filename, + "current": cur, + "total": t.total, + }) +} + +type SharePlugin struct { + sidechannel transport.SidechannelOptions + DownloadDir string + cfg config.ShareConfig + TLSConfig *tls.Config + Logger *zap.Logger + bus *events.Bus +} + +func NewSharePlugin(downloadDir string, cfg config.ShareConfig, tlsConfig *tls.Config, bus *events.Bus, logger *zap.Logger, options ...transport.SidechannelOptions) *SharePlugin { + var sidechannel transport.SidechannelOptions + if len(options) > 0 { + sidechannel = options[0] + } + return &SharePlugin{ + sidechannel: sidechannel, + DownloadDir: downloadDir, + cfg: cfg, + TLSConfig: tlsConfig, + Logger: logger.With(zap.String("plugin", "share")), + bus: bus, + } +} + +// ShareBody includes the missing Android metadata (LastModified/CreationTime) +type ShareBody struct { + Filename string `json:"filename"` + NumberOfFiles int `json:"numberOfFiles,omitempty"` + TotalPayloadSize int64 `json:"totalPayloadSize,omitempty"` + LastModified int64 `json:"lastModified,omitempty"` + CreationTime int64 `json:"creationTime,omitempty"` + Text string `json:"text,omitempty"` + Url string `json:"url,omitempty"` +} + +func (p *SharePlugin) Name() string { return "Share" } + +func (p *SharePlugin) Timeout() time.Duration { return 0 } + +func (p *SharePlugin) IncomingTypes() []string { + return []string{"kdeconnect.share.request"} +} + +func (p *SharePlugin) OutgoingTypes() []string { + return []string{"kdeconnect.share.request"} +} diff --git a/internal/plugins/sms/attachment.go b/internal/plugins/sms/attachment.go new file mode 100644 index 0000000..3dad891 --- /dev/null +++ b/internal/plugins/sms/attachment.go @@ -0,0 +1,126 @@ +package sms + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net" + "os" + "path/filepath" + "strings" + "unicode" + + "github.com/bethropolis/kcd/internal/cert" + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +// handleAttachmentFile downloads an MMS attachment file sent by the phone. +func (p *SMSPlugin) handleAttachmentFile(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + if pkt.Body == nil { + return nil + } + + var body AttachmentFileBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return fmt.Errorf("sms: unmarshal attachment body: %w", err) + } + + if body.Filename == "" { + return nil + } + + if pkt.PayloadTransferInfo == nil || pkt.PayloadTransferInfo.Port == 0 { + p.logger.Warn("sms: attachment file received without side-channel transfer info (inline payload not supported)") + return nil + } + + cleanName := cleanFilename(body.Filename) + destPath := filepath.Join(p.cacheDir, cleanName) + + remoteIP := dev.RemoteIP() + if remoteIP == nil { + return fmt.Errorf("sms: failed to resolve remote peer IP") + } + + port := pkt.PayloadTransferInfo.Port + payloadSize := pkt.PayloadSize + expectedFP := cert.PinnedFingerprint(dev.PeerCert()) + + go func() { + if err := p.receiveAttachment(ctx, remoteIP, port, payloadSize, destPath, expectedFP); err != nil { + p.logger.Error("sms: attachment download failed", zap.Error(err)) + return + } + p.logger.Info("sms: attachment downloaded", + zap.String("path", destPath), + zap.String("filename", body.Filename), + ) + if p.bus != nil { + p.bus.Publish(events.TypeSMSAttachment, dev.ID(), map[string]any{ + "filename": body.Filename, + "path": destPath, + "thread_id": body.ThreadID, + }) + } + }() + + return nil +} + +// receiveAttachment connects to the phone's side-channel port and downloads +// the attachment file over TLS. The stream is capped at the declared +// payload size (itself bounded by maxSMSAttachmentBytes) so a malicious +// peer can't fill the disk with an unbounded stream. +func (p *SMSPlugin) receiveAttachment(ctx context.Context, ip net.IP, port int, size int64, destPath string, expectedFP string) error { + if size <= 0 || size > maxSMSAttachmentBytes { + return fmt.Errorf("sms: refusing attachment with invalid size %d (limit %d)", size, maxSMSAttachmentBytes) + } + conn, err := transport.DialSidechannel(ctx, ip, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) + if err != nil { + return fmt.Errorf("sms: connect to attachment side-channel: %w", err) + } + defer conn.Close() + + f, err := os.Create(destPath) + if err != nil { + return fmt.Errorf("sms: create attachment file: %w", err) + } + defer f.Close() + + _, err = io.Copy(f, io.LimitReader(conn, size)) + if err != nil { + os.Remove(destPath) // don't leave a corrupt partial behind + return fmt.Errorf("sms: receive attachment data: %w", err) + } + + return nil +} + +// maxFilenameLength caps attachment filenames to keep them manageable. +const maxFilenameLength = 128 + +// cleanFilename strips path components to prevent directory traversal in +// attachment file paths. Backslashes are normalized first (Windows-style +// paths), control characters dropped, and overlong names truncated. +func cleanFilename(name string) string { + name = strings.ReplaceAll(name, "\\", "/") + name = filepath.Base(name) + name = strings.Map(func(r rune) rune { + if unicode.IsControl(r) { + return -1 + } + return r + }, name) + if len(name) > maxFilenameLength { + name = strings.ToValidUTF8(name[:maxFilenameLength], "") + } + if name == "." || name == ".." || name == "/" || name == "" { + return "downloaded_attachment" + } + return name +} diff --git a/internal/plugins/sms/handle.go b/internal/plugins/sms/handle.go new file mode 100644 index 0000000..c635910 --- /dev/null +++ b/internal/plugins/sms/handle.go @@ -0,0 +1,95 @@ +package sms + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/plugin" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +// --- Handle ---------------------------------------------------------------- + +func (p *SMSPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + switch pkt.Type { + case PacketTypeSMSMessages: + return p.handleMessages(ctx, dev, pkt) + case PacketTypeSMSAttachmentFile: + return p.handleAttachmentFile(ctx, dev, pkt) + } + return nil +} + +// handleMessages parses a batch of SMS messages from the phone and publishes +// one event per message. +func (p *SMSPlugin) handleMessages(_ context.Context, dev device.Sender, pkt *protocol.Packet) error { + if pkt.Body == nil { + return nil + } + + var batch SMSMessagesPacket + if err := json.Unmarshal(pkt.Body, &batch); err != nil { + return fmt.Errorf("sms: unmarshal messages batch: %w", err) + } + + if len(batch.Messages) > maxSMSMessages { + return fmt.Errorf("sms: messages batch too large: %d (max %d)", len(batch.Messages), maxSMSMessages) + } + + for _, msg := range batch.Messages { + if msg.Body == "" { + continue + } + + msg := msg // capture + + sender := "" + if len(msg.Addresses) > 0 { + sender = msg.Addresses[0].Address + } + + p.logger.Debug("sms: message received", + zap.String("from", sender), + zap.String("body", msg.Body), + zap.Int64("thread_id", msg.ThreadID), + ) + + if p.bus != nil { + payload := map[string]any{ + "body": msg.Body, + "sender": sender, + "date": msg.Date, + "type": msg.Type, + "thread_id": msg.ThreadID, + "read": bool(msg.Read), + "event": msg.Event, + "u_id": msg.UID, + "sub_id": msg.SubID, + } + if len(msg.Attachments) > 0 { + payload["attachments"] = msg.Attachments + } + p.bus.Publish(events.TypeSMSIncoming, dev.ID(), payload) + } + + if p.cfg.NotifyIncoming { + msgText := msg.Body + if len(msgText) > 120 { + msgText = msgText[:120] + "…" + } + title := fmt.Sprintf("SMS from %s", sender) + plugin.RunCommandAsync(p.logger, "notify-send", + "-a", p.notifications.AppName(), + "-i", "dialog-information", + title, + msgText, + ) + } + } + + return nil +} diff --git a/internal/plugins/sms/send.go b/internal/plugins/sms/send.go new file mode 100644 index 0000000..e4f0388 --- /dev/null +++ b/internal/plugins/sms/send.go @@ -0,0 +1,75 @@ +package sms + +import ( + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/protocol" +) + +// --- SMS sending ----------------------------------------------------------- + +func (p *SMSPlugin) SendSMS(dev device.Sender, phoneNumber, message string) error { + // v2 schema: the phone reads only messageBody, with addresses as the + // primary recipient list (phoneNumber stays as a legacy fallback for + // older peers). Without addresses/version the phone sends a blank SMS. + body := map[string]any{ + "version": 2, + "addresses": []map[string]string{{"address": phoneNumber}}, + "messageBody": message, + "phoneNumber": phoneNumber, + } + pkt, err := protocol.NewPacket(PacketTypeSMSRequest, body) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// --- Conversation browsing (Phase 2) --------------------------------------- + +// RequestConversations asks the phone for a summary of all conversations. +// Bodyless requests use an empty object (never null) on the wire; see +// contacts.RequestSync for why explicit null is dangerous. +func (p *SMSPlugin) RequestConversations(dev device.Sender) error { + pkt, err := protocol.NewPacket(PacketTypeSMSRequestConvs, map[string]any{}) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// RequestConversation asks the phone for messages in a specific thread. +// Pass -1 for rangeStartTimestamp or numberToRequest for no limit. +func (p *SMSPlugin) RequestConversation(dev device.Sender, threadID int64, rangeStartTimestamp int64, numberToRequest int64) error { + body := map[string]any{ + "threadID": threadID, + } + if rangeStartTimestamp >= 0 { + body["rangeStartTimestamp"] = rangeStartTimestamp + } + if numberToRequest >= 0 { + body["numberToRequest"] = numberToRequest + } + pkt, err := protocol.NewPacket(PacketTypeSMSRequestConv, body) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// RequestAttachment asks the phone to send an MMS attachment file. +func (p *SMSPlugin) RequestAttachment(dev device.Sender, partID int64, uniqueIdentifier string) error { + body := map[string]any{ + "part_id": partID, + "unique_identifier": uniqueIdentifier, + } + pkt, err := protocol.NewPacket(PacketTypeSMSRequestAtt, body) + if err != nil { + return err + } + return dev.Send(pkt) +} + +// --- Lifecycle ------------------------------------------------------------- + +func (p *SMSPlugin) OnConnect(dev device.Sender) {} +func (p *SMSPlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/sms/sms.go b/internal/plugins/sms/sms.go deleted file mode 100644 index a6956a8..0000000 --- a/internal/plugins/sms/sms.go +++ /dev/null @@ -1,390 +0,0 @@ -package sms - -import ( - "context" - "crypto/tls" - "encoding/json" - "fmt" - "io" - "net" - "os" - "path/filepath" - "strings" - "time" - "unicode" - - "github.com/bethropolis/kcd/internal/cert" - "github.com/bethropolis/kcd/internal/config" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/plugin" - "github.com/bethropolis/kcd/internal/protocol" - "github.com/bethropolis/kcd/internal/transport" - "go.uber.org/zap" -) - -const ( - PacketTypeSMSMessages = "kdeconnect.sms.messages" - PacketTypeSMSRequest = "kdeconnect.sms.request" - PacketTypeSMSRequestConvs = "kdeconnect.sms.request_conversations" - PacketTypeSMSRequestConv = "kdeconnect.sms.request_conversation" - PacketTypeSMSRequestAtt = "kdeconnect.sms.request_attachment" - PacketTypeSMSAttachmentFile = "kdeconnect.sms.attachment_file" - - maxSMSMessages = 1000 // safety limit to prevent OOM from malicious payload - - // maxSMSAttachmentBytes caps a single MMS attachment download (mirrors - // the clipboard 50MB safety limit). Anything larger is refused before - // any bytes hit disk. - maxSMSAttachmentBytes = 50 * 1024 * 1024 -) - -// SMSPlugin implements SMS sending, receiving, conversation browsing, and MMS -// attachment handling for KDE Connect. -type SMSPlugin struct { - sidechannel transport.SidechannelOptions - notifications config.NotificationConfig - cfg config.SMSConfig - bus *events.Bus - tlsConfig *tls.Config - logger *zap.Logger - cacheDir string -} - -// Options customizes storage, network timeouts, and desktop notification identity. -type Options struct { - CacheDir string - Sidechannel transport.SidechannelOptions - Notifications config.NotificationConfig -} - -func NewSMSPlugin(cfg config.SMSConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...Options) *SMSPlugin { - var opts Options - if len(options) > 0 { - opts = options[0] - } - cacheDir := opts.CacheDir - if cacheDir == "" { - cacheDir = filepath.Join(os.TempDir(), "kcd", "sms-attachments") - } - _ = os.MkdirAll(cacheDir, 0700) - - return &SMSPlugin{ - sidechannel: opts.Sidechannel, - notifications: opts.Notifications, - cfg: cfg, - bus: bus, - tlsConfig: tlsConfig, - logger: logger.With(zap.String("plugin", "sms")), - cacheDir: cacheDir, - } -} - -func (p *SMSPlugin) Name() string { return "SMS" } -func (p *SMSPlugin) Timeout() time.Duration { return 5 * time.Second } - -func (p *SMSPlugin) IncomingTypes() []string { - return []string{PacketTypeSMSMessages, PacketTypeSMSAttachmentFile} -} - -func (p *SMSPlugin) OutgoingTypes() []string { - return []string{ - PacketTypeSMSRequest, - PacketTypeSMSRequestConvs, - PacketTypeSMSRequestConv, - PacketTypeSMSRequestAtt, - } -} - -// --- Packet body types ------------------------------------------------------- - -type SMSMessagesPacket struct { - Version int `json:"version"` - Messages []SMSMessage `json:"messages"` -} - -type SMSMessage struct { - Event int `json:"event"` - Body string `json:"body"` - Addresses []SMSAddress `json:"addresses"` - Date int64 `json:"date"` - Type int `json:"type"` - ThreadID int64 `json:"thread_id"` - Read protocol.FlexBool `json:"read"` - UID int64 `json:"u_id,omitempty"` - SubID int `json:"sub_id,omitempty"` - Attachments []SMSAttachment `json:"attachments,omitempty"` -} - -type SMSAddress struct { - Address string `json:"address"` -} - -type SMSAttachment struct { - PartID int64 `json:"part_id"` - MimeType string `json:"mime_type"` - EncodedThumbnail string `json:"encoded_thumbnail,omitempty"` - UniqueIdentifier string `json:"unique_identifier"` -} - -type AttachmentFileBody struct { - Filename string `json:"filename"` - ThreadID int64 `json:"thread_id,omitempty"` -} - -// --- Handle ---------------------------------------------------------------- - -func (p *SMSPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - switch pkt.Type { - case PacketTypeSMSMessages: - return p.handleMessages(ctx, dev, pkt) - case PacketTypeSMSAttachmentFile: - return p.handleAttachmentFile(ctx, dev, pkt) - } - return nil -} - -// handleMessages parses a batch of SMS messages from the phone and publishes -// one event per message. -func (p *SMSPlugin) handleMessages(_ context.Context, dev device.Sender, pkt *protocol.Packet) error { - if pkt.Body == nil { - return nil - } - - var batch SMSMessagesPacket - if err := json.Unmarshal(pkt.Body, &batch); err != nil { - return fmt.Errorf("sms: unmarshal messages batch: %w", err) - } - - if len(batch.Messages) > maxSMSMessages { - return fmt.Errorf("sms: messages batch too large: %d (max %d)", len(batch.Messages), maxSMSMessages) - } - - for _, msg := range batch.Messages { - if msg.Body == "" { - continue - } - - msg := msg // capture - - sender := "" - if len(msg.Addresses) > 0 { - sender = msg.Addresses[0].Address - } - - p.logger.Debug("sms: message received", - zap.String("from", sender), - zap.String("body", msg.Body), - zap.Int64("thread_id", msg.ThreadID), - ) - - if p.bus != nil { - payload := map[string]any{ - "body": msg.Body, - "sender": sender, - "date": msg.Date, - "type": msg.Type, - "thread_id": msg.ThreadID, - "read": bool(msg.Read), - "event": msg.Event, - "u_id": msg.UID, - "sub_id": msg.SubID, - } - if len(msg.Attachments) > 0 { - payload["attachments"] = msg.Attachments - } - p.bus.Publish(events.TypeSMSIncoming, dev.ID(), payload) - } - - if p.cfg.NotifyIncoming { - msgText := msg.Body - if len(msgText) > 120 { - msgText = msgText[:120] + "…" - } - title := fmt.Sprintf("SMS from %s", sender) - plugin.RunCommandAsync(p.logger, "notify-send", - "-a", p.notifications.AppName(), - "-i", "dialog-information", - title, - msgText, - ) - } - } - - return nil -} - -// handleAttachmentFile downloads an MMS attachment file sent by the phone. -func (p *SMSPlugin) handleAttachmentFile(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - if pkt.Body == nil { - return nil - } - - var body AttachmentFileBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return fmt.Errorf("sms: unmarshal attachment body: %w", err) - } - - if body.Filename == "" { - return nil - } - - if pkt.PayloadTransferInfo == nil || pkt.PayloadTransferInfo.Port == 0 { - p.logger.Warn("sms: attachment file received without side-channel transfer info (inline payload not supported)") - return nil - } - - cleanName := cleanFilename(body.Filename) - destPath := filepath.Join(p.cacheDir, cleanName) - - remoteIP := dev.RemoteIP() - if remoteIP == nil { - return fmt.Errorf("sms: failed to resolve remote peer IP") - } - - port := pkt.PayloadTransferInfo.Port - payloadSize := pkt.PayloadSize - expectedFP := cert.PinnedFingerprint(dev.PeerCert()) - - go func() { - if err := p.receiveAttachment(ctx, remoteIP, port, payloadSize, destPath, expectedFP); err != nil { - p.logger.Error("sms: attachment download failed", zap.Error(err)) - return - } - p.logger.Info("sms: attachment downloaded", - zap.String("path", destPath), - zap.String("filename", body.Filename), - ) - if p.bus != nil { - p.bus.Publish(events.TypeSMSAttachment, dev.ID(), map[string]any{ - "filename": body.Filename, - "path": destPath, - "thread_id": body.ThreadID, - }) - } - }() - - return nil -} - -// receiveAttachment connects to the phone's side-channel port and downloads -// the attachment file over TLS. The stream is capped at the declared -// payload size (itself bounded by maxSMSAttachmentBytes) so a malicious -// peer can't fill the disk with an unbounded stream. -func (p *SMSPlugin) receiveAttachment(ctx context.Context, ip net.IP, port int, size int64, destPath string, expectedFP string) error { - if size <= 0 || size > maxSMSAttachmentBytes { - return fmt.Errorf("sms: refusing attachment with invalid size %d (limit %d)", size, maxSMSAttachmentBytes) - } - conn, err := transport.DialSidechannel(ctx, ip, port, p.tlsConfig, expectedFP, p.logger, p.sidechannel) - if err != nil { - return fmt.Errorf("sms: connect to attachment side-channel: %w", err) - } - defer conn.Close() - - f, err := os.Create(destPath) - if err != nil { - return fmt.Errorf("sms: create attachment file: %w", err) - } - defer f.Close() - - _, err = io.Copy(f, io.LimitReader(conn, size)) - if err != nil { - os.Remove(destPath) // don't leave a corrupt partial behind - return fmt.Errorf("sms: receive attachment data: %w", err) - } - - return nil -} - -// --- SMS sending ----------------------------------------------------------- - -func (p *SMSPlugin) SendSMS(dev device.Sender, phoneNumber, message string) error { - // v2 schema: the phone reads only messageBody, with addresses as the - // primary recipient list (phoneNumber stays as a legacy fallback for - // older peers). Without addresses/version the phone sends a blank SMS. - body := map[string]any{ - "version": 2, - "addresses": []map[string]string{{"address": phoneNumber}}, - "messageBody": message, - "phoneNumber": phoneNumber, - } - pkt, err := protocol.NewPacket(PacketTypeSMSRequest, body) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// --- Conversation browsing (Phase 2) --------------------------------------- - -// RequestConversations asks the phone for a summary of all conversations. -// Bodyless requests use an empty object (never null) on the wire; see -// contacts.RequestSync for why explicit null is dangerous. -func (p *SMSPlugin) RequestConversations(dev device.Sender) error { - pkt, err := protocol.NewPacket(PacketTypeSMSRequestConvs, map[string]any{}) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// RequestConversation asks the phone for messages in a specific thread. -// Pass -1 for rangeStartTimestamp or numberToRequest for no limit. -func (p *SMSPlugin) RequestConversation(dev device.Sender, threadID int64, rangeStartTimestamp int64, numberToRequest int64) error { - body := map[string]any{ - "threadID": threadID, - } - if rangeStartTimestamp >= 0 { - body["rangeStartTimestamp"] = rangeStartTimestamp - } - if numberToRequest >= 0 { - body["numberToRequest"] = numberToRequest - } - pkt, err := protocol.NewPacket(PacketTypeSMSRequestConv, body) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// RequestAttachment asks the phone to send an MMS attachment file. -func (p *SMSPlugin) RequestAttachment(dev device.Sender, partID int64, uniqueIdentifier string) error { - body := map[string]any{ - "part_id": partID, - "unique_identifier": uniqueIdentifier, - } - pkt, err := protocol.NewPacket(PacketTypeSMSRequestAtt, body) - if err != nil { - return err - } - return dev.Send(pkt) -} - -// maxFilenameLength caps attachment filenames to keep them manageable. -const maxFilenameLength = 128 - -// cleanFilename strips path components to prevent directory traversal in -// attachment file paths. Backslashes are normalized first (Windows-style -// paths), control characters dropped, and overlong names truncated. -func cleanFilename(name string) string { - name = strings.ReplaceAll(name, "\\", "/") - name = filepath.Base(name) - name = strings.Map(func(r rune) rune { - if unicode.IsControl(r) { - return -1 - } - return r - }, name) - if len(name) > maxFilenameLength { - name = strings.ToValidUTF8(name[:maxFilenameLength], "") - } - if name == "." || name == ".." || name == "/" || name == "" { - return "downloaded_attachment" - } - return name -} - -// --- Lifecycle ------------------------------------------------------------- - -func (p *SMSPlugin) OnConnect(dev device.Sender) {} -func (p *SMSPlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/sms/types.go b/internal/plugins/sms/types.go new file mode 100644 index 0000000..08b111d --- /dev/null +++ b/internal/plugins/sms/types.go @@ -0,0 +1,123 @@ +package sms + +import ( + "crypto/tls" + "os" + "path/filepath" + "time" + + "github.com/bethropolis/kcd/internal/config" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "github.com/bethropolis/kcd/internal/transport" + "go.uber.org/zap" +) + +const ( + PacketTypeSMSMessages = "kdeconnect.sms.messages" + PacketTypeSMSRequest = "kdeconnect.sms.request" + PacketTypeSMSRequestConvs = "kdeconnect.sms.request_conversations" + PacketTypeSMSRequestConv = "kdeconnect.sms.request_conversation" + PacketTypeSMSRequestAtt = "kdeconnect.sms.request_attachment" + PacketTypeSMSAttachmentFile = "kdeconnect.sms.attachment_file" + + maxSMSMessages = 1000 // safety limit to prevent OOM from malicious payload + + // maxSMSAttachmentBytes caps a single MMS attachment download (mirrors + // the clipboard 50MB safety limit). Anything larger is refused before + // any bytes hit disk. + maxSMSAttachmentBytes = 50 * 1024 * 1024 +) + +// SMSPlugin implements SMS sending, receiving, conversation browsing, and MMS +// attachment handling for KDE Connect. +type SMSPlugin struct { + sidechannel transport.SidechannelOptions + notifications config.NotificationConfig + cfg config.SMSConfig + bus *events.Bus + tlsConfig *tls.Config + logger *zap.Logger + cacheDir string +} + +// Options customizes storage, network timeouts, and desktop notification identity. +type Options struct { + CacheDir string + Sidechannel transport.SidechannelOptions + Notifications config.NotificationConfig +} + +func NewSMSPlugin(cfg config.SMSConfig, bus *events.Bus, tlsConfig *tls.Config, logger *zap.Logger, options ...Options) *SMSPlugin { + var opts Options + if len(options) > 0 { + opts = options[0] + } + cacheDir := opts.CacheDir + if cacheDir == "" { + cacheDir = filepath.Join(os.TempDir(), "kcd", "sms-attachments") + } + _ = os.MkdirAll(cacheDir, 0700) + + return &SMSPlugin{ + sidechannel: opts.Sidechannel, + notifications: opts.Notifications, + cfg: cfg, + bus: bus, + tlsConfig: tlsConfig, + logger: logger.With(zap.String("plugin", "sms")), + cacheDir: cacheDir, + } +} + +func (p *SMSPlugin) Name() string { return "SMS" } +func (p *SMSPlugin) Timeout() time.Duration { return 5 * time.Second } + +func (p *SMSPlugin) IncomingTypes() []string { + return []string{PacketTypeSMSMessages, PacketTypeSMSAttachmentFile} +} + +func (p *SMSPlugin) OutgoingTypes() []string { + return []string{ + PacketTypeSMSRequest, + PacketTypeSMSRequestConvs, + PacketTypeSMSRequestConv, + PacketTypeSMSRequestAtt, + } +} + +// --- Packet body types ------------------------------------------------------- + +type SMSMessagesPacket struct { + Version int `json:"version"` + Messages []SMSMessage `json:"messages"` +} + +type SMSMessage struct { + Event int `json:"event"` + Body string `json:"body"` + Addresses []SMSAddress `json:"addresses"` + Date int64 `json:"date"` + Type int `json:"type"` + ThreadID int64 `json:"thread_id"` + Read protocol.FlexBool `json:"read"` + UID int64 `json:"u_id,omitempty"` + SubID int `json:"sub_id,omitempty"` + Attachments []SMSAttachment `json:"attachments,omitempty"` +} + +type SMSAddress struct { + Address string `json:"address"` +} + +type SMSAttachment struct { + PartID int64 `json:"part_id"` + MimeType string `json:"mime_type"` + EncodedThumbnail string `json:"encoded_thumbnail,omitempty"` + UniqueIdentifier string `json:"unique_identifier"` +} + +type AttachmentFileBody struct { + Filename string `json:"filename"` + ThreadID int64 `json:"thread_id,omitempty"` +} diff --git a/internal/plugins/systemvolume/backend.go b/internal/plugins/systemvolume/backend.go new file mode 100644 index 0000000..40d0fbc --- /dev/null +++ b/internal/plugins/systemvolume/backend.go @@ -0,0 +1,110 @@ +package systemvolume + +import ( + "context" + "os/exec" + "strconv" + "strings" + "time" +) + +// getSinks returns a list of available audio output sinks. +func (p *SystemVolumePlugin) getSinks() []SinkInfo { + switch p.backend { + case "wpctl": + return p.getSinksWpctl() + case "pactl": + return p.getSinksPactl() + } + return nil +} + +func (p *SystemVolumePlugin) getSinksWpctl() []SinkInfo { + // Get current volume from wpctl: wpctl get-volume @DEFAULT_AUDIO_SINK@ + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + out, err := exec.CommandContext(ctx, "wpctl", "get-volume", "@DEFAULT_AUDIO_SINK@").Output() + if err != nil { + return nil + } + // Output: "Volume: 0.75 [MUTED]" or "Volume: 0.75" + line := strings.TrimSpace(string(out)) + muted := strings.Contains(line, "[MUTED]") + line = strings.ReplaceAll(line, "[MUTED]", "") + parts := strings.Fields(line) + vol := 75 + if len(parts) >= 2 { + if f, err := strconv.ParseFloat(parts[1], 64); err == nil { + vol = int(f * 100) + } + } + return []SinkInfo{{ + Name: "@DEFAULT_AUDIO_SINK@", + Description: "Default Output", + Volume: vol, + Muted: muted, + MaxVolume: 100, + }} +} + +func (p *SystemVolumePlugin) getSinksPactl() []SinkInfo { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + out, err := exec.CommandContext(ctx, "pactl", "get-sink-volume", "@DEFAULT_SINK@").Output() + if err != nil { + return nil + } + // Very rough parse: look for the first percentage + vol := 75 + for _, field := range strings.Fields(string(out)) { + if strings.HasSuffix(field, "%") { + if v, err := strconv.Atoi(strings.TrimSuffix(field, "%")); err == nil { + vol = v + break + } + } + } + muteCtx, muteCancel := context.WithTimeout(context.Background(), 5*time.Second) + defer muteCancel() + muteOut, _ := exec.CommandContext(muteCtx, "pactl", "get-sink-mute", "@DEFAULT_SINK@").Output() + muted := strings.Contains(string(muteOut), "yes") + return []SinkInfo{{ + Name: "@DEFAULT_SINK@", + Description: "Default Output", + Volume: vol, + Muted: muted, + MaxVolume: 100, + }} +} + +// setVolume applies volume and mute settings via the detected backend. +func (p *SystemVolumePlugin) setVolume(name string, volume int, muted bool) error { + return p.setVolumeStr(name, strconv.Itoa(volume), muted) +} + +func (p *SystemVolumePlugin) setVolumeStr(_ string, volumeStr string, muted bool) error { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + switch p.backend { + case "wpctl": + pct := volumeStr + "%" + if err := exec.CommandContext(ctx, "wpctl", "set-volume", "@DEFAULT_AUDIO_SINK@", pct).Run(); err != nil { + return err + } + muteArg := "0" + if muted { + muteArg = "1" + } + return exec.CommandContext(ctx, "wpctl", "set-mute", "@DEFAULT_AUDIO_SINK@", muteArg).Run() + case "pactl": + if err := exec.CommandContext(ctx, "pactl", "set-sink-volume", "@DEFAULT_SINK@", volumeStr+"%").Run(); err != nil { + return err + } + muteArg := "false" + if muted { + muteArg = "true" + } + return exec.CommandContext(ctx, "pactl", "set-sink-mute", "@DEFAULT_SINK@", muteArg).Run() + } + return nil +} diff --git a/internal/plugins/systemvolume/handle.go b/internal/plugins/systemvolume/handle.go new file mode 100644 index 0000000..4e4d78b --- /dev/null +++ b/internal/plugins/systemvolume/handle.go @@ -0,0 +1,73 @@ +package systemvolume + +import ( + "context" + "encoding/json" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" +) + +func (p *SystemVolumePlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + if p.backend == "" { + return nil // No audio backend, silently ignore. + } + + var body VolumeBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + return err + } + + // Phone is requesting the list of audio sinks. + if body.RequestSinks { + go func() { + sinks := p.getSinks() + pkt, err := protocol.NewPacket("kdeconnect.systemvolume", sinkListBody{SinkList: sinks}) + if err != nil { + p.logger.Error("systemvolume: failed to create sink list packet", zap.Error(err)) + return + } + if err := dev.Send(pkt); err != nil { + p.logger.Error("systemvolume: failed to send sink list", zap.Error(err)) + } + }() + return nil + } + + // Phone is setting volume or mute. + go func() { + if body.Name == "" { + body.Name = "@DEFAULT_AUDIO_SINK@" + } + if err := p.setVolume(body.Name, body.Volume, body.Muted); err != nil { + p.logger.Warn("systemvolume: failed to set volume", zap.Error(err)) + return + } + if p.bus != nil { + p.bus.Publish(events.TypeVolumeUpdate, dev.ID(), map[string]any{ + "name": body.Name, + "volume": body.Volume, + "muted": body.Muted, + }) + } + }() + + return nil +} + +func (p *SystemVolumePlugin) OnConnect(dev device.Sender) { + if p.backend == "" { + return + } + go func() { + sinks := p.getSinks() + pkt, err := protocol.NewPacket("kdeconnect.systemvolume", sinkListBody{SinkList: sinks}) + if err != nil { + return + } + _ = dev.Send(pkt) + }() +} +func (p *SystemVolumePlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/systemvolume/systemvolume.go b/internal/plugins/systemvolume/systemvolume.go deleted file mode 100644 index 0a581bf..0000000 --- a/internal/plugins/systemvolume/systemvolume.go +++ /dev/null @@ -1,232 +0,0 @@ -// Package systemvolume implements the KDE Connect System Volume plugin. -// It allows the phone to query and control the PC's audio volume. -package systemvolume - -import ( - "context" - "encoding/json" - "os/exec" - "strconv" - "strings" - "time" - - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" - "github.com/bethropolis/kcd/internal/protocol" - "go.uber.org/zap" -) - -// SystemVolumePlugin handles volume control packets from the phone. -type SystemVolumePlugin struct { - logger *zap.Logger - bus *events.Bus - backend string // "wpctl" or "pactl" -} - -func NewSystemVolumePlugin(bus *events.Bus, logger *zap.Logger) *SystemVolumePlugin { - p := &SystemVolumePlugin{ - logger: logger.With(zap.String("plugin", "systemvolume")), - bus: bus, - } - // Detect available audio backend at init time. - if _, err := exec.LookPath("wpctl"); err == nil { - p.backend = "wpctl" - } else if _, err := exec.LookPath("pactl"); err == nil { - p.backend = "pactl" - } else { - p.logger.Warn("systemvolume: no audio backend found (wpctl or pactl required)") - } - return p -} - -type VolumeBody struct { - RequestSinks bool `json:"requestSinks,omitempty"` - Name string `json:"name,omitempty"` - Volume int `json:"volume,omitempty"` - Muted bool `json:"muted,omitempty"` - MaxVolume int `json:"maxVolume,omitempty"` -} - -type SinkInfo struct { - Name string `json:"name"` - Description string `json:"description"` - Volume int `json:"volume"` - Muted bool `json:"muted"` - MaxVolume int `json:"maxVolume"` -} - -func (p *SystemVolumePlugin) Name() string { return "SystemVolume" } -func (p *SystemVolumePlugin) Timeout() time.Duration { return 5 * time.Second } -func (p *SystemVolumePlugin) IncomingTypes() []string { - return []string{"kdeconnect.systemvolume.request"} -} -func (p *SystemVolumePlugin) OutgoingTypes() []string { return []string{"kdeconnect.systemvolume"} } - -func (p *SystemVolumePlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { - if p.backend == "" { - return nil // No audio backend, silently ignore. - } - - var body VolumeBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return err - } - - // Phone is requesting the list of audio sinks. - if body.RequestSinks { - go func() { - sinks := p.getSinks() - type sinkListBody struct { - SinkList []SinkInfo `json:"sinkList"` - } - pkt, err := protocol.NewPacket("kdeconnect.systemvolume", sinkListBody{SinkList: sinks}) - if err != nil { - p.logger.Error("systemvolume: failed to create sink list packet", zap.Error(err)) - return - } - if err := dev.Send(pkt); err != nil { - p.logger.Error("systemvolume: failed to send sink list", zap.Error(err)) - } - }() - return nil - } - - // Phone is setting volume or mute. - go func() { - if body.Name == "" { - body.Name = "@DEFAULT_AUDIO_SINK@" - } - if err := p.setVolume(body.Name, body.Volume, body.Muted); err != nil { - p.logger.Warn("systemvolume: failed to set volume", zap.Error(err)) - return - } - if p.bus != nil { - p.bus.Publish(events.TypeVolumeUpdate, dev.ID(), map[string]any{ - "name": body.Name, - "volume": body.Volume, - "muted": body.Muted, - }) - } - }() - - return nil -} - -// getSinks returns a list of available audio output sinks. -func (p *SystemVolumePlugin) getSinks() []SinkInfo { - switch p.backend { - case "wpctl": - return p.getSinksWpctl() - case "pactl": - return p.getSinksPactl() - } - return nil -} - -func (p *SystemVolumePlugin) getSinksWpctl() []SinkInfo { - // Get current volume from wpctl: wpctl get-volume @DEFAULT_AUDIO_SINK@ - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - out, err := exec.CommandContext(ctx, "wpctl", "get-volume", "@DEFAULT_AUDIO_SINK@").Output() - if err != nil { - return nil - } - // Output: "Volume: 0.75 [MUTED]" or "Volume: 0.75" - line := strings.TrimSpace(string(out)) - muted := strings.Contains(line, "[MUTED]") - line = strings.ReplaceAll(line, "[MUTED]", "") - parts := strings.Fields(line) - vol := 75 - if len(parts) >= 2 { - if f, err := strconv.ParseFloat(parts[1], 64); err == nil { - vol = int(f * 100) - } - } - return []SinkInfo{{ - Name: "@DEFAULT_AUDIO_SINK@", - Description: "Default Output", - Volume: vol, - Muted: muted, - MaxVolume: 100, - }} -} - -func (p *SystemVolumePlugin) getSinksPactl() []SinkInfo { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - out, err := exec.CommandContext(ctx, "pactl", "get-sink-volume", "@DEFAULT_SINK@").Output() - if err != nil { - return nil - } - // Very rough parse: look for the first percentage - vol := 75 - for _, field := range strings.Fields(string(out)) { - if strings.HasSuffix(field, "%") { - if v, err := strconv.Atoi(strings.TrimSuffix(field, "%")); err == nil { - vol = v - break - } - } - } - muteCtx, muteCancel := context.WithTimeout(context.Background(), 5*time.Second) - defer muteCancel() - muteOut, _ := exec.CommandContext(muteCtx, "pactl", "get-sink-mute", "@DEFAULT_SINK@").Output() - muted := strings.Contains(string(muteOut), "yes") - return []SinkInfo{{ - Name: "@DEFAULT_SINK@", - Description: "Default Output", - Volume: vol, - Muted: muted, - MaxVolume: 100, - }} -} - -// setVolume applies volume and mute settings via the detected backend. -func (p *SystemVolumePlugin) setVolume(name string, volume int, muted bool) error { - return p.setVolumeStr(name, strconv.Itoa(volume), muted) -} - -func (p *SystemVolumePlugin) setVolumeStr(_ string, volumeStr string, muted bool) error { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - switch p.backend { - case "wpctl": - pct := volumeStr + "%" - if err := exec.CommandContext(ctx, "wpctl", "set-volume", "@DEFAULT_AUDIO_SINK@", pct).Run(); err != nil { - return err - } - muteArg := "0" - if muted { - muteArg = "1" - } - return exec.CommandContext(ctx, "wpctl", "set-mute", "@DEFAULT_AUDIO_SINK@", muteArg).Run() - case "pactl": - if err := exec.CommandContext(ctx, "pactl", "set-sink-volume", "@DEFAULT_SINK@", volumeStr+"%").Run(); err != nil { - return err - } - muteArg := "false" - if muted { - muteArg = "true" - } - return exec.CommandContext(ctx, "pactl", "set-sink-mute", "@DEFAULT_SINK@", muteArg).Run() - } - return nil -} - -func (p *SystemVolumePlugin) OnConnect(dev device.Sender) { - if p.backend == "" { - return - } - go func() { - sinks := p.getSinks() - type sinkListBody struct { - SinkList []SinkInfo `json:"sinkList"` - } - pkt, err := protocol.NewPacket("kdeconnect.systemvolume", sinkListBody{SinkList: sinks}) - if err != nil { - return - } - _ = dev.Send(pkt) - }() -} -func (p *SystemVolumePlugin) OnDisconnect(dev device.Sender) {} diff --git a/internal/plugins/systemvolume/types.go b/internal/plugins/systemvolume/types.go new file mode 100644 index 0000000..1467b2a --- /dev/null +++ b/internal/plugins/systemvolume/types.go @@ -0,0 +1,62 @@ +// Package systemvolume implements the KDE Connect System Volume plugin. +// It allows the phone to query and control the PC's audio volume. +package systemvolume + +import ( + "os/exec" + "time" + + "github.com/bethropolis/kcd/internal/events" + "go.uber.org/zap" +) + +// SystemVolumePlugin handles volume control packets from the phone. +type SystemVolumePlugin struct { + logger *zap.Logger + bus *events.Bus + backend string // "wpctl" or "pactl" +} + +func NewSystemVolumePlugin(bus *events.Bus, logger *zap.Logger) *SystemVolumePlugin { + p := &SystemVolumePlugin{ + logger: logger.With(zap.String("plugin", "systemvolume")), + bus: bus, + } + // Detect available audio backend at init time. + if _, err := exec.LookPath("wpctl"); err == nil { + p.backend = "wpctl" + } else if _, err := exec.LookPath("pactl"); err == nil { + p.backend = "pactl" + } else { + p.logger.Warn("systemvolume: no audio backend found (wpctl or pactl required)") + } + return p +} + +type VolumeBody struct { + RequestSinks bool `json:"requestSinks,omitempty"` + Name string `json:"name,omitempty"` + Volume int `json:"volume,omitempty"` + Muted bool `json:"muted,omitempty"` + MaxVolume int `json:"maxVolume,omitempty"` +} + +type SinkInfo struct { + Name string `json:"name"` + Description string `json:"description"` + Volume int `json:"volume"` + Muted bool `json:"muted"` + MaxVolume int `json:"maxVolume"` +} + +// sinkListBody is the outbound sink list payload. +type sinkListBody struct { + SinkList []SinkInfo `json:"sinkList"` +} + +func (p *SystemVolumePlugin) Name() string { return "SystemVolume" } +func (p *SystemVolumePlugin) Timeout() time.Duration { return 5 * time.Second } +func (p *SystemVolumePlugin) IncomingTypes() []string { + return []string{"kdeconnect.systemvolume.request"} +} +func (p *SystemVolumePlugin) OutgoingTypes() []string { return []string{"kdeconnect.systemvolume"} } diff --git a/pkg/client/client.go b/pkg/client/client.go index 2d8a5ea..d1381d0 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -9,10 +9,7 @@ import ( "net" "time" - "github.com/bethropolis/kcd/internal/device" - "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/ipc" - "github.com/bethropolis/kcd/internal/plugins/contacts" ) // Client connects to the kcd daemon via Unix socket. @@ -78,465 +75,3 @@ func (c *Client) Call(cmd string, payload interface{}) (*ipc.Response, error) { return &res, nil } - -// Connect requests the daemon to manually connect to a device by IP. -func (c *Client) Connect(ip string) error { - _, err := c.Call(ipc.CmdConnect, ipc.ConnectPayload{IP: ip}) - return err -} - -// Devices queries the daemon for all known devices. -func (c *Client) Devices() ([]device.DeviceInfo, error) { - res, err := c.Call(ipc.CmdDevices, nil) - if err != nil { - return nil, err - } - - var devices []device.DeviceInfo - if err := json.Unmarshal(res.Data, &devices); err != nil { - return nil, fmt.Errorf("decode devices: %w", err) - } - return devices, nil -} - -// PairListen enters listen mode: waits for an incoming pair request, auto-accepts -// it, and returns the paired device info. Blocks up to 60 seconds. -func (c *Client) PairListen() (*ipc.PairListenResult, error) { - // Copy the client so a long listen never changes concurrent calls' deadlines. - listenClient := *c - listenClient.Timeout = c.PairListenTimeout - if listenClient.Timeout <= 0 { - listenClient.Timeout = 70 * time.Second - } - - resp, err := listenClient.Call(ipc.CmdPairListen, nil) - if err != nil { - return nil, err - } - var result ipc.PairListenResult - if len(resp.Data) > 0 { - _ = json.Unmarshal(resp.Data, &result) - } - return &result, nil -} - -// Pair requests the daemon to pair with a specific device. -func (c *Client) Pair(deviceID string) error { - _, err := c.Call(ipc.CmdPair, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Unpair requests the daemon to unpair and forget a specific device. -func (c *Client) Unpair(deviceID string) error { - _, err := c.Call(ipc.CmdUnpair, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Ping sends a ping packet to a specific device. -func (c *Client) Ping(deviceID string) error { - _, err := c.Call(ipc.CmdPing, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Battery queries the daemon for a device's battery state. -func (c *Client) Battery(deviceID string) (int, bool, error) { - res, err := c.Call(ipc.CmdBattery, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return 0, false, err - } - - var data struct { - Charge int `json:"charge"` - Charging bool `json:"charging"` - } - if err := json.Unmarshal(res.Data, &data); err != nil { - return 0, false, err - } - return data.Charge, data.Charging, nil -} - -// ClipboardPush triggers an outgoing clipboard sync from desktop to device. -func (c *Client) ClipboardPush(deviceID string) error { - _, err := c.Call(ipc.CmdClipboardPush, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Connectivity returns the last raw connectivity report for a device. -// Callers decode it (same shape as connectivity.update event payloads). -func (c *Client) Connectivity(deviceID string) (json.RawMessage, error) { - res, err := c.Call(ipc.CmdConnectivity, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return nil, err - } - return res.Data, nil -} - -// RunList requests the remote device to send its command list. -func (c *Client) RunList(deviceID string) error { - _, err := c.Call(ipc.CmdRunList, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// RunExec requests the remote device to execute a specific command key. -func (c *Client) RunExec(deviceID string, key string) error { - _, err := c.Call(ipc.CmdRunExec, ipc.DevicePayload{DeviceID: deviceID, Key: key}) - return err -} - -// ShareFile requests the daemon to send a local file to the remote device. -func (c *Client) ShareFile(deviceID string, filePath string) error { - _, err := c.Call(ipc.CmdShare, ipc.SharePayload{DeviceID: deviceID, FilePath: filePath}) - return err -} - -// SftpMount requests the daemon to initiate an SFTP connection to the remote device. -func (c *Client) SftpMount(deviceID string) error { - _, err := c.Call(ipc.CmdSftpMount, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// SftpInfo returns the cached SFTP connection details for a device. -func (c *Client) SftpInfo(deviceID string) (*ipc.SftpInfoResponse, error) { - resp, err := c.Call(ipc.CmdSftpInfo, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return nil, err - } - var info ipc.SftpInfoResponse - if len(resp.Data) > 0 { - _ = json.Unmarshal(resp.Data, &info) - } - return &info, nil -} - -// SftpVolumes returns the list of available storage volumes from a device. -func (c *Client) SftpVolumes(deviceID string) ([]ipc.StorageVolumeResponse, error) { - resp, err := c.Call(ipc.CmdSftpVolumes, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return nil, err - } - var volumes []ipc.StorageVolumeResponse - if len(resp.Data) > 0 { - _ = json.Unmarshal(resp.Data, &volumes) - } - return volumes, nil -} - -// SftpMountLocal requests the daemon to request SFTP credentials from the -// phone, wait for the response, mount via sshfs, and open the result in -// the default file manager. Returns the local browse path on success. -func (c *Client) SftpMountLocal(deviceID string) (string, error) { - resp, err := c.Call(ipc.CmdSftpMountLocal, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return "", err - } - var result struct { - Path string `json:"path"` - } - if len(resp.Data) > 0 { - _ = json.Unmarshal(resp.Data, &result) - } - return result.Path, nil -} - -// SftpUnmount cleanly unmounts a previously mounted phone filesystem. -func (c *Client) SftpUnmount(deviceID string) error { - _, err := c.Call(ipc.CmdSftpUnmount, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// SftpBrowse requests fresh SFTP credentials from the phone and either lists -// available volumes (volume arg empty) or mounts the specified volume. -// volume can be an index (0-based), volume name, or path. -// Returns the mount path (empty if listing) and available volumes. -func (c *Client) SftpBrowse(deviceID string, volume string) (string, []ipc.StorageVolumeResponse, error) { - resp, err := c.Call(ipc.CmdSftpBrowse, ipc.SftpBrowsePayload{ - DeviceID: deviceID, - Volume: volume, - }) - if err != nil { - return "", nil, err - } - var result ipc.SftpBrowseResponse - if len(resp.Data) > 0 { - _ = json.Unmarshal(resp.Data, &result) - } - return result.Path, result.Volumes, nil -} - -// BroadcastStart asks the daemon to begin UDP/mDNS broadcasting. -func (c *Client) BroadcastStart() error { - _, err := c.Call(ipc.CmdBroadcastStart, nil) - return err -} - -// BroadcastStop asks the daemon to stop UDP/mDNS broadcasting. -func (c *Client) BroadcastStop() error { - _, err := c.Call(ipc.CmdBroadcastStop, nil) - return err -} - -// Status returns runtime status information from the daemon. -func (c *Client) Status() (*ipc.StatusResponse, error) { - res, err := c.Call(ipc.CmdStatus, nil) - if err != nil { - return nil, err - } - var resp ipc.StatusResponse - if err := json.Unmarshal(res.Data, &resp); err != nil { - return nil, err - } - return &resp, nil -} - -// NotifyReply requests the daemon to send a reply to an Android notification. -func (c *Client) NotifyReply(deviceID, replyID, message string) error { - _, err := c.Call(ipc.CmdNotifyReply, ipc.NotifyReplyPayload{ - DeviceID: deviceID, - ReplyID: replyID, - Message: message, - }) - return err -} - -// NotifyDismiss requests the daemon to clear a notification on the remote device. -func (c *Client) NotifyDismiss(deviceID, notificationID string) error { - _, err := c.Call(ipc.CmdNotifyDismiss, ipc.NotifyDismissPayload{ - DeviceID: deviceID, - NotificationID: notificationID, - }) - return err -} - -// CallMute requests the daemon to mute an incoming call on the remote device. -func (c *Client) CallMute(deviceID string) error { - _, err := c.Call(ipc.CmdCallMute, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// FindMyPhone requests the remote device to ring loudly. -func (c *Client) FindMyPhone(deviceID string) error { - _, err := c.Call(ipc.CmdFindMyPhone, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Lock requests the daemon to lock the session. -func (c *Client) Lock(deviceID string) error { - _, err := c.Call(ipc.CmdLock, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// Unlock requests the daemon to unlock the session. -func (c *Client) Unlock(deviceID string) error { - _, err := c.Call(ipc.CmdUnlock, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// SendSMS requests the remote device to send an SMS. -func (c *Client) SendSMS(deviceID, phoneNumber, message string) error { - _, err := c.Call(ipc.CmdSendSMS, ipc.SMSPayload{ - DeviceID: deviceID, - PhoneNumber: phoneNumber, - Message: message, - }) - return err -} - -// SmsRequestConversations asks a device to send a list of all SMS conversations. -func (c *Client) SmsRequestConversations(deviceID string) error { - _, err := c.Call(ipc.CmdSmsRequestConvs, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// SmsRequestConversation asks a device to send messages from a specific thread. -func (c *Client) SmsRequestConversation(deviceID string, threadID int64) error { - _, err := c.Call(ipc.CmdSmsRequestConv, ipc.SMSConvPayload{DeviceID: deviceID, ThreadID: threadID}) - return err -} - -// SmsRequestAttachment asks a device to send an MMS attachment file. -func (c *Client) SmsRequestAttachment(deviceID string, partID int64, uniqueIdentifier string) error { - _, err := c.Call(ipc.CmdSmsRequestAttachment, ipc.SMSAttachmentPayload{ - DeviceID: deviceID, - PartID: partID, - UniqueIdentifier: uniqueIdentifier, - }) - return err -} - -// ContactsSync asks the daemon to start a contacts sync round with a device. -// Results arrive async via the contacts.updated event. -func (c *Client) ContactsSync(deviceID string) error { - _, err := c.Call(ipc.CmdContactsSync, ipc.DevicePayload{DeviceID: deviceID}) - return err -} - -// ContactsList returns cached contact summaries for a device (empty when -// never synced — absent means unknown). -func (c *Client) ContactsList(deviceID string) ([]contacts.ContactSummary, error) { - res, err := c.Call(ipc.CmdContactsList, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return nil, err - } - var list []contacts.ContactSummary - if err := json.Unmarshal(res.Data, &list); err != nil { - return nil, fmt.Errorf("decode contacts: %w", err) - } - return list, nil -} - -// MprisStatus returns MPRIS plugin debug information. -func (c *Client) MprisStatus() (*ipc.MprisStatusResponse, error) { - res, err := c.Call(ipc.CmdMprisStatus, nil) - if err != nil { - return nil, err - } - var resp ipc.MprisStatusResponse - if err := json.Unmarshal(res.Data, &resp); err != nil { - return nil, fmt.Errorf("decode mpris status: %w", err) - } - return &resp, nil -} - -// MprisAction sends a media control action to a remote device. -// deviceID may be empty to auto-select the first connected device. -func (c *Client) MprisAction(deviceID, player, action string) error { - _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ - DeviceID: deviceID, - Player: player, - Action: action, - }) - return err -} - -// MprisVolume sends a volume change to a remote device's player. -func (c *Client) MprisVolume(deviceID, player string, volume int) error { - v := volume - _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ - DeviceID: deviceID, - Player: player, - Volume: &v, - }) - return err -} - -// MprisSeek sends a seek command to a remote device's player. -func (c *Client) MprisSeek(deviceID, player string, seek int64) error { - s := seek - _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ - DeviceID: deviceID, - Player: player, - Seek: &s, - }) - return err -} - -// MprisRemote returns the list of remote MPRIS players with their current state. -func (c *Client) MprisRemote() (*ipc.MprisRemoteResponse, error) { - res, err := c.Call(ipc.CmdMprisRemote, nil) - if err != nil { - return nil, err - } - var resp ipc.MprisRemoteResponse - if err := json.Unmarshal(res.Data, &resp); err != nil { - return nil, fmt.Errorf("decode mpris remote: %w", err) - } - return &resp, nil -} - -// RemoteVolumeList returns the last known sink list for a device. -func (c *Client) RemoteVolumeList(deviceID string) ([]byte, error) { - resp, err := c.Call(ipc.CmdRemoteVolumeList, ipc.DevicePayload{DeviceID: deviceID}) - if err != nil { - return nil, err - } - return resp.Data, nil -} - -// RemoteVolumeSet sets the volume for a specific sink on a remote device. -func (c *Client) RemoteVolumeSet(deviceID, sinkName string, volume int) error { - _, err := c.Call(ipc.CmdRemoteVolumeSet, struct { - DeviceID string `json:"deviceId"` - Name string `json:"name"` - Volume int `json:"volume"` - }{DeviceID: deviceID, Name: sinkName, Volume: volume}) - return err -} - -// RemoteVolumeMute sets the mute state for a specific sink on a remote device. -func (c *Client) RemoteVolumeMute(deviceID, sinkName string, muted bool) error { - _, err := c.Call(ipc.CmdRemoteVolumeMute, struct { - DeviceID string `json:"deviceId"` - Name string `json:"name"` - Muted bool `json:"muted"` - }{DeviceID: deviceID, Name: sinkName, Muted: muted}) - return err -} - -// Watch subscribes to daemon events and streams them to the given channel. -func (c *Client) Watch(ctx context.Context, filter []string, ch chan<- events.Event) error { - dialer := net.Dialer{} - conn, err := dialer.DialContext(ctx, "unix", c.SocketPath) - if err != nil { - return fmt.Errorf("kcd daemon not running or socket error: %w", err) - } - - payload, _ := json.Marshal(ipc.WatchPayload{Events: filter}) - req := ipc.Request{ - Command: ipc.CmdWatch, - Payload: payload, - } - - reqBytes, _ := json.Marshal(req) - reqBytes = append(reqBytes, '\n') - - if _, err := conn.Write(reqBytes); err != nil { - conn.Close() - return fmt.Errorf("write request: %w", err) - } - - reader := bufio.NewReader(conn) - resBytes, err := reader.ReadBytes('\n') - if err != nil { - conn.Close() - return fmt.Errorf("read response: %w", err) - } - - var res ipc.Response - if err := json.Unmarshal(resBytes, &res); err != nil { - conn.Close() - return fmt.Errorf("unmarshal response: %w", err) - } - - if !res.OK { - conn.Close() - return fmt.Errorf("daemon error: %s", res.Error) - } - - defer conn.Close() - - // Create a goroutine to close the connection if context is canceled - go func() { - <-ctx.Done() - conn.Close() - }() - - for { - line, err := reader.ReadBytes('\n') - if err != nil { - if ctx.Err() != nil { - return ctx.Err() - } - return fmt.Errorf("stream read error: %w", err) - } - - var ev events.Event - if err := json.Unmarshal(line, &ev); err != nil { - continue - } - select { - case ch <- ev: - case <-ctx.Done(): - return ctx.Err() - } - } -} diff --git a/pkg/client/client_actions.go b/pkg/client/client_actions.go new file mode 100644 index 0000000..d73e14a --- /dev/null +++ b/pkg/client/client_actions.go @@ -0,0 +1,97 @@ +package client + +import ( + "encoding/json" + + "github.com/bethropolis/kcd/internal/ipc" +) + +// Ping sends a ping packet to a specific device. +func (c *Client) Ping(deviceID string) error { + _, err := c.Call(ipc.CmdPing, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// Battery queries the daemon for a device's battery state. +func (c *Client) Battery(deviceID string) (int, bool, error) { + res, err := c.Call(ipc.CmdBattery, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return 0, false, err + } + + var data struct { + Charge int `json:"charge"` + Charging bool `json:"charging"` + } + if err := json.Unmarshal(res.Data, &data); err != nil { + return 0, false, err + } + return data.Charge, data.Charging, nil +} + +// ClipboardPush triggers an outgoing clipboard sync from desktop to device. +func (c *Client) ClipboardPush(deviceID string) error { + _, err := c.Call(ipc.CmdClipboardPush, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// RunList requests the remote device to send its command list. +func (c *Client) RunList(deviceID string) error { + _, err := c.Call(ipc.CmdRunList, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// RunExec requests the remote device to execute a specific command key. +func (c *Client) RunExec(deviceID string, key string) error { + _, err := c.Call(ipc.CmdRunExec, ipc.DevicePayload{DeviceID: deviceID, Key: key}) + return err +} + +// ShareFile requests the daemon to send a local file to the remote device. +func (c *Client) ShareFile(deviceID string, filePath string) error { + _, err := c.Call(ipc.CmdShare, ipc.SharePayload{DeviceID: deviceID, FilePath: filePath}) + return err +} + +// NotifyReply requests the daemon to send a reply to an Android notification. +func (c *Client) NotifyReply(deviceID, replyID, message string) error { + _, err := c.Call(ipc.CmdNotifyReply, ipc.NotifyReplyPayload{ + DeviceID: deviceID, + ReplyID: replyID, + Message: message, + }) + return err +} + +// NotifyDismiss requests the daemon to clear a notification on the remote device. +func (c *Client) NotifyDismiss(deviceID, notificationID string) error { + _, err := c.Call(ipc.CmdNotifyDismiss, ipc.NotifyDismissPayload{ + DeviceID: deviceID, + NotificationID: notificationID, + }) + return err +} + +// CallMute requests the daemon to mute an incoming call on the remote device. +func (c *Client) CallMute(deviceID string) error { + _, err := c.Call(ipc.CmdCallMute, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// FindMyPhone requests the remote device to ring loudly. +func (c *Client) FindMyPhone(deviceID string) error { + _, err := c.Call(ipc.CmdFindMyPhone, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// Lock requests the daemon to lock the session. +func (c *Client) Lock(deviceID string) error { + _, err := c.Call(ipc.CmdLock, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// Unlock requests the daemon to unlock the session. +func (c *Client) Unlock(deviceID string) error { + _, err := c.Call(ipc.CmdUnlock, ipc.DevicePayload{DeviceID: deviceID}) + return err +} diff --git a/pkg/client/client_devices.go b/pkg/client/client_devices.go new file mode 100644 index 0000000..d0e04eb --- /dev/null +++ b/pkg/client/client_devices.go @@ -0,0 +1,98 @@ +package client + +import ( + "encoding/json" + "fmt" + "time" + + "github.com/bethropolis/kcd/internal/device" + "github.com/bethropolis/kcd/internal/ipc" +) + +// Connect requests the daemon to manually connect to a device by IP. +func (c *Client) Connect(ip string) error { + _, err := c.Call(ipc.CmdConnect, ipc.ConnectPayload{IP: ip}) + return err +} + +// Devices queries the daemon for all known devices. +func (c *Client) Devices() ([]device.DeviceInfo, error) { + res, err := c.Call(ipc.CmdDevices, nil) + if err != nil { + return nil, err + } + + var devices []device.DeviceInfo + if err := json.Unmarshal(res.Data, &devices); err != nil { + return nil, fmt.Errorf("decode devices: %w", err) + } + return devices, nil +} + +// PairListen enters listen mode: waits for an incoming pair request, auto-accepts +// it, and returns the paired device info. Blocks up to 60 seconds. +func (c *Client) PairListen() (*ipc.PairListenResult, error) { + // Copy the client so a long listen never changes concurrent calls' deadlines. + listenClient := *c + listenClient.Timeout = c.PairListenTimeout + if listenClient.Timeout <= 0 { + listenClient.Timeout = 70 * time.Second + } + + resp, err := listenClient.Call(ipc.CmdPairListen, nil) + if err != nil { + return nil, err + } + var result ipc.PairListenResult + if len(resp.Data) > 0 { + _ = json.Unmarshal(resp.Data, &result) + } + return &result, nil +} + +// Pair requests the daemon to pair with a specific device. +func (c *Client) Pair(deviceID string) error { + _, err := c.Call(ipc.CmdPair, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// Unpair requests the daemon to unpair and forget a specific device. +func (c *Client) Unpair(deviceID string) error { + _, err := c.Call(ipc.CmdUnpair, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// BroadcastStart asks the daemon to begin UDP/mDNS broadcasting. +func (c *Client) BroadcastStart() error { + _, err := c.Call(ipc.CmdBroadcastStart, nil) + return err +} + +// BroadcastStop asks the daemon to stop UDP/mDNS broadcasting. +func (c *Client) BroadcastStop() error { + _, err := c.Call(ipc.CmdBroadcastStop, nil) + return err +} + +// Status returns runtime status information from the daemon. +func (c *Client) Status() (*ipc.StatusResponse, error) { + res, err := c.Call(ipc.CmdStatus, nil) + if err != nil { + return nil, err + } + var resp ipc.StatusResponse + if err := json.Unmarshal(res.Data, &resp); err != nil { + return nil, err + } + return &resp, nil +} + +// Connectivity returns the last raw connectivity report for a device. +// Callers decode it (same shape as connectivity.update event payloads). +func (c *Client) Connectivity(deviceID string) (json.RawMessage, error) { + res, err := c.Call(ipc.CmdConnectivity, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return nil, err + } + return res.Data, nil +} diff --git a/pkg/client/client_mpris.go b/pkg/client/client_mpris.go new file mode 100644 index 0000000..707c273 --- /dev/null +++ b/pkg/client/client_mpris.go @@ -0,0 +1,96 @@ +package client + +import ( + "encoding/json" + "fmt" + + "github.com/bethropolis/kcd/internal/ipc" +) + +// MprisStatus returns MPRIS plugin debug information. +func (c *Client) MprisStatus() (*ipc.MprisStatusResponse, error) { + res, err := c.Call(ipc.CmdMprisStatus, nil) + if err != nil { + return nil, err + } + var resp ipc.MprisStatusResponse + if err := json.Unmarshal(res.Data, &resp); err != nil { + return nil, fmt.Errorf("decode mpris status: %w", err) + } + return &resp, nil +} + +// MprisAction sends a media control action to a remote device. +// deviceID may be empty to auto-select the first connected device. +func (c *Client) MprisAction(deviceID, player, action string) error { + _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ + DeviceID: deviceID, + Player: player, + Action: action, + }) + return err +} + +// MprisVolume sends a volume change to a remote device's player. +func (c *Client) MprisVolume(deviceID, player string, volume int) error { + v := volume + _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ + DeviceID: deviceID, + Player: player, + Volume: &v, + }) + return err +} + +// MprisSeek sends a seek command to a remote device's player. +func (c *Client) MprisSeek(deviceID, player string, seek int64) error { + s := seek + _, err := c.Call(ipc.CmdMprisAction, ipc.MprisActionPayload{ + DeviceID: deviceID, + Player: player, + Seek: &s, + }) + return err +} + +// MprisRemote returns the list of remote MPRIS players with their current state. +func (c *Client) MprisRemote() (*ipc.MprisRemoteResponse, error) { + res, err := c.Call(ipc.CmdMprisRemote, nil) + if err != nil { + return nil, err + } + var resp ipc.MprisRemoteResponse + if err := json.Unmarshal(res.Data, &resp); err != nil { + return nil, fmt.Errorf("decode mpris remote: %w", err) + } + return &resp, nil +} + +// RemoteVolumeList returns the last known sink list for a device. +func (c *Client) RemoteVolumeList(deviceID string) ([]byte, error) { + resp, err := c.Call(ipc.CmdRemoteVolumeList, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return nil, err + } + return resp.Data, nil +} + +// RemoteVolumeSet sets the volume for a specific sink on a remote device. +func (c *Client) RemoteVolumeSet(deviceID, sinkName string, volume int) error { + _, err := c.Call(ipc.CmdRemoteVolumeSet, struct { + DeviceID string `json:"deviceId"` + Name string `json:"name"` + Volume int `json:"volume"` + }{DeviceID: deviceID, Name: sinkName, Volume: volume}) + return err +} + +// RemoteVolumeMute sets the mute state for a specific sink on a remote device. +func (c *Client) RemoteVolumeMute(deviceID, sinkName string, muted bool) error { + _, err := c.Call(ipc.CmdRemoteVolumeMute, struct { + DeviceID string `json:"deviceId"` + Name string `json:"name"` + Muted bool `json:"muted"` + }{DeviceID: deviceID, Name: sinkName, Muted: muted}) + return err +} diff --git a/pkg/client/client_sftp.go b/pkg/client/client_sftp.go new file mode 100644 index 0000000..60f7bd2 --- /dev/null +++ b/pkg/client/client_sftp.go @@ -0,0 +1,81 @@ +package client + +import ( + "encoding/json" + + "github.com/bethropolis/kcd/internal/ipc" +) + +// SftpMount requests the daemon to initiate an SFTP connection to the remote device. +func (c *Client) SftpMount(deviceID string) error { + _, err := c.Call(ipc.CmdSftpMount, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// SftpInfo returns the cached SFTP connection details for a device. +func (c *Client) SftpInfo(deviceID string) (*ipc.SftpInfoResponse, error) { + resp, err := c.Call(ipc.CmdSftpInfo, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return nil, err + } + var info ipc.SftpInfoResponse + if len(resp.Data) > 0 { + _ = json.Unmarshal(resp.Data, &info) + } + return &info, nil +} + +// SftpVolumes returns the list of available storage volumes from a device. +func (c *Client) SftpVolumes(deviceID string) ([]ipc.StorageVolumeResponse, error) { + resp, err := c.Call(ipc.CmdSftpVolumes, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return nil, err + } + var volumes []ipc.StorageVolumeResponse + if len(resp.Data) > 0 { + _ = json.Unmarshal(resp.Data, &volumes) + } + return volumes, nil +} + +// SftpMountLocal requests the daemon to request SFTP credentials from the +// phone, wait for the response, mount via sshfs, and open the result in +// the default file manager. Returns the local browse path on success. +func (c *Client) SftpMountLocal(deviceID string) (string, error) { + resp, err := c.Call(ipc.CmdSftpMountLocal, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return "", err + } + var result struct { + Path string `json:"path"` + } + if len(resp.Data) > 0 { + _ = json.Unmarshal(resp.Data, &result) + } + return result.Path, nil +} + +// SftpUnmount cleanly unmounts a previously mounted phone filesystem. +func (c *Client) SftpUnmount(deviceID string) error { + _, err := c.Call(ipc.CmdSftpUnmount, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// SftpBrowse requests fresh SFTP credentials from the phone and either lists +// available volumes (volume arg empty) or mounts the specified volume. +// volume can be an index (0-based), volume name, or path. +// Returns the mount path (empty if listing) and available volumes. +func (c *Client) SftpBrowse(deviceID string, volume string) (string, []ipc.StorageVolumeResponse, error) { + resp, err := c.Call(ipc.CmdSftpBrowse, ipc.SftpBrowsePayload{ + DeviceID: deviceID, + Volume: volume, + }) + if err != nil { + return "", nil, err + } + var result ipc.SftpBrowseResponse + if len(resp.Data) > 0 { + _ = json.Unmarshal(resp.Data, &result) + } + return result.Path, result.Volumes, nil +} diff --git a/pkg/client/client_sms_contacts.go b/pkg/client/client_sms_contacts.go new file mode 100644 index 0000000..cd6066c --- /dev/null +++ b/pkg/client/client_sms_contacts.go @@ -0,0 +1,62 @@ +package client + +import ( + "encoding/json" + "fmt" + + "github.com/bethropolis/kcd/internal/ipc" + "github.com/bethropolis/kcd/internal/plugins/contacts" +) + +// SendSMS requests the remote device to send an SMS. +func (c *Client) SendSMS(deviceID, phoneNumber, message string) error { + _, err := c.Call(ipc.CmdSendSMS, ipc.SMSPayload{ + DeviceID: deviceID, + PhoneNumber: phoneNumber, + Message: message, + }) + return err +} + +// SmsRequestConversations asks a device to send a list of all SMS conversations. +func (c *Client) SmsRequestConversations(deviceID string) error { + _, err := c.Call(ipc.CmdSmsRequestConvs, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// SmsRequestConversation asks a device to send messages from a specific thread. +func (c *Client) SmsRequestConversation(deviceID string, threadID int64) error { + _, err := c.Call(ipc.CmdSmsRequestConv, ipc.SMSConvPayload{DeviceID: deviceID, ThreadID: threadID}) + return err +} + +// SmsRequestAttachment asks a device to send an MMS attachment file. +func (c *Client) SmsRequestAttachment(deviceID string, partID int64, uniqueIdentifier string) error { + _, err := c.Call(ipc.CmdSmsRequestAttachment, ipc.SMSAttachmentPayload{ + DeviceID: deviceID, + PartID: partID, + UniqueIdentifier: uniqueIdentifier, + }) + return err +} + +// ContactsSync asks the daemon to start a contacts sync round with a device. +// Results arrive async via the contacts.updated event. +func (c *Client) ContactsSync(deviceID string) error { + _, err := c.Call(ipc.CmdContactsSync, ipc.DevicePayload{DeviceID: deviceID}) + return err +} + +// ContactsList returns cached contact summaries for a device (empty when +// never synced — absent means unknown). +func (c *Client) ContactsList(deviceID string) ([]contacts.ContactSummary, error) { + res, err := c.Call(ipc.CmdContactsList, ipc.DevicePayload{DeviceID: deviceID}) + if err != nil { + return nil, err + } + var list []contacts.ContactSummary + if err := json.Unmarshal(res.Data, &list); err != nil { + return nil, fmt.Errorf("decode contacts: %w", err) + } + return list, nil +} diff --git a/pkg/client/client_watch.go b/pkg/client/client_watch.go new file mode 100644 index 0000000..5f75016 --- /dev/null +++ b/pkg/client/client_watch.go @@ -0,0 +1,81 @@ +package client + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "net" + + "github.com/bethropolis/kcd/internal/events" + "github.com/bethropolis/kcd/internal/ipc" +) + +// Watch subscribes to daemon events and streams them to the given channel. +func (c *Client) Watch(ctx context.Context, filter []string, ch chan<- events.Event) error { + dialer := net.Dialer{} + conn, err := dialer.DialContext(ctx, "unix", c.SocketPath) + if err != nil { + return fmt.Errorf("kcd daemon not running or socket error: %w", err) + } + + payload, _ := json.Marshal(ipc.WatchPayload{Events: filter}) + req := ipc.Request{ + Command: ipc.CmdWatch, + Payload: payload, + } + + reqBytes, _ := json.Marshal(req) + reqBytes = append(reqBytes, '\n') + + if _, err := conn.Write(reqBytes); err != nil { + conn.Close() + return fmt.Errorf("write request: %w", err) + } + + reader := bufio.NewReader(conn) + resBytes, err := reader.ReadBytes('\n') + if err != nil { + conn.Close() + return fmt.Errorf("read response: %w", err) + } + + var res ipc.Response + if err := json.Unmarshal(resBytes, &res); err != nil { + conn.Close() + return fmt.Errorf("unmarshal response: %w", err) + } + + if !res.OK { + conn.Close() + return fmt.Errorf("daemon error: %s", res.Error) + } + + defer conn.Close() + + // Create a goroutine to close the connection if context is canceled + go func() { + <-ctx.Done() + conn.Close() + }() + + for { + line, err := reader.ReadBytes('\n') + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return fmt.Errorf("stream read error: %w", err) + } + + var ev events.Event + if err := json.Unmarshal(line, &ev); err != nil { + continue + } + select { + case ch <- ev: + case <-ctx.Done(): + return ctx.Err() + } + } +} From d498d853a2a93f41e202095fb8a6cd13369bad8c Mon Sep 17 00:00:00 2001 From: bethropolis <66518866+bethropolis@users.noreply.github.com> Date: Fri, 18 Sep 2026 20:41:02 +0300 Subject: [PATCH 12/12] fix(ci): satisfy noctx lint and detach telephony test logger - sidechannel_test: use ListenConfig.Listen instead of net.Listen. - telephony missed-call test: plugin notifies via a background goroutine that can outlive the test; give it a detached Nop logger so late writes cannot panic the test runner. --- internal/plugins/telephony/telephony_test.go | 5 ++++- internal/transport/sidechannel_test.go | 4 ++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/internal/plugins/telephony/telephony_test.go b/internal/plugins/telephony/telephony_test.go index a5c3e50..24fc42a 100644 --- a/internal/plugins/telephony/telephony_test.go +++ b/internal/plugins/telephony/telephony_test.go @@ -11,6 +11,7 @@ import ( "github.com/bethropolis/kcd/internal/device" "github.com/bethropolis/kcd/internal/events" "github.com/bethropolis/kcd/internal/protocol" + "go.uber.org/zap" "go.uber.org/zap/zaptest" ) @@ -63,7 +64,9 @@ func TestHandleStringIsCancel(t *testing.T) { func TestHandleMissedCallEvent(t *testing.T) { logger := zaptest.NewLogger(t) bus := events.NewBus(logger) - p := NewTelephonyPlugin(bus, logger) + // The plugin notifies via a background goroutine that can outlive the + // test; a test-bound logger would panic on late writes, so detach it. + p := NewTelephonyPlugin(bus, zap.NewNop()) sub := bus.Subscribe(0, events.TypeTelephonyMissed) defer sub.Close() diff --git a/internal/transport/sidechannel_test.go b/internal/transport/sidechannel_test.go index 517dec9..3652f1e 100644 --- a/internal/transport/sidechannel_test.go +++ b/internal/transport/sidechannel_test.go @@ -21,7 +21,7 @@ import ( // It reports the peer certificate so the test can pin the fingerprint. func startSidechannelServer(t *testing.T, tlsConfig *tls.Config, payload []byte, delay time.Duration) (addr string, fp string, done <-chan error) { t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") + ln, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } @@ -137,7 +137,7 @@ func TestDialSidechannelPinMismatch(t *testing.T) { // handshake must surface an error within roughly the configured budget. func TestDialSidechannelSetupTimeout(t *testing.T) { // Plain TCP listener: accepts but never speaks TLS. - ln, err := net.Listen("tcp", "127.0.0.1:0") + ln, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) }