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_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/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/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/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..f99803e 100644 --- a/cmd/kcd/main.go +++ b/cmd/kcd/main.go @@ -7,7 +7,6 @@ import ( "log" "os" "os/signal" - "strings" "syscall" "time" @@ -31,11 +30,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 @@ -120,6 +129,7 @@ func main() { watchCmd, sftpCmd, replyCmd, + dismissCmd, callCmd, findmyphoneCmd, lockCmd, @@ -184,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 12a1efd..1db81cb 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"`, `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 | +| `[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. @@ -96,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, ... --- @@ -543,6 +589,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/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/config/config.go b/internal/config/config.go index dbf2da8..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" ) @@ -27,6 +24,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 +66,9 @@ func Defaults() *Config { c.TCPPort = 1716 c.LogLevel = "info" + 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() c.Commands = make(map[string]string) c.CommandsPerDevice = make(map[string]map[string]string) @@ -100,155 +104,8 @@ func Load(path string) (*Config, error) { } cfg.ConfigPath = path - return cfg, nil -} - -// Validate checks required fields and returns an error if any are invalid. -func (c *Config) Validate() error { - 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) - } + if err := cfg.Validate(); err != nil { + return nil, 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 per-app notification filtering. -type NotificationConfig map[string]string - -// 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 + return cfg, 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_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..3139633 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,227 @@ +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", "60s"}) { + 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" +transfer_idle_timeout = "90s" +[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", "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" { + 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/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/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..d3b0495 --- /dev/null +++ b/internal/config/runtime.go @@ -0,0 +1,77 @@ +package config + +import ( + "fmt" + "time" +) + +// 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"` + TransferIdleTimeout string `toml:"transfer_idle_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}, + {"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}, + {"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..5b6632e 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,14 @@ 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) + // 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} @@ -112,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{ @@ -125,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/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/daemon/plugins.go b/internal/daemon/plugins.go index 9eaedf9..e375ebc 100644 --- a/internal/daemon/plugins.go +++ b/internal/daemon/plugins.go @@ -28,38 +28,43 @@ 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), + IdleTimeout: config.Duration(cfg.Network.TransferIdleTimeout), + } 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 +88,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..fa42f37 100644 --- a/internal/daemon/transport.go +++ b/internal/daemon/transport.go @@ -9,7 +9,7 @@ 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" "github.com/bethropolis/kcd/internal/plugin" @@ -18,103 +18,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"). -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) { - 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: 5 * time.Second, - } - 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 - } - 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, 10*time.Second) - 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); 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 @@ -138,9 +41,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 +159,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 +170,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 +190,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 +202,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 +262,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 +271,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() } @@ -378,208 +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) 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() >= reconnectFlapThreshold { - 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) - } - - 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, -) { - const maxBackoff = 5 * time.Minute - 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) - 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 := 1716 - if p := dev.LastPort(); validDialPort(p) { - port = p - } - DialDevice(ctx, ip, port, dev.ID(), protocol.ProtocolVersion, identity, cfg, devices, plugins, localDeviceID, logger, false) - - 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++ - } -} 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/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") + } + }) + } +} 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 deleted file mode 100644 index c352ead..0000000 --- a/internal/device/device.go +++ /dev/null @@ -1,617 +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() { - d.pairDialRequested.Store(true) - d.pairIntentUntil.Store(time.Now().Add(pairDialIntentTTL).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/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/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 8b79fb8..0000000 --- a/internal/discovery/discovery.go +++ /dev/null @@ -1,369 +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 - 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) *BroadcasterController { - return &BroadcasterController{ - identityPacket: identity, - interval: interval, - 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, - 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 - 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 := 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/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/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/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..78fc722 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 @@ -155,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"} } @@ -188,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) @@ -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/ipc/proto.go b/internal/ipc/proto.go index e0adfb2..cbbe19b 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"` @@ -121,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. 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/battery.go deleted file mode 100644 index d3e787e..0000000 --- a/internal/plugins/battery/battery.go +++ /dev/null @@ -1,211 +0,0 @@ -package battery - -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" - "github.com/bethropolis/kcd/internal/protocol" - "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 { - 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 { - return &BatteryPlugin{ - cfg: cfg, - bus: bus, - logger: logger.With(zap.String("plugin", "battery")), - } -} - -// BatteryBody represents the body of a kdeconnect.battery packet. -type BatteryBody struct { - CurrentCharge int `json:"currentCharge"` - IsCharging bool `json:"isCharging"` - ThresholdEvent int `json:"thresholdEvent"` -} - -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 { - case "kdeconnect.battery": - var body BatteryBody - if err := json.Unmarshal(pkt.Body, &body); err != nil { - return fmt.Errorf("battery: decode body: %w", err) - } - - dev.UpdateBattery(body.CurrentCharge, body.IsCharging) - - if body.ThresholdEvent != thresholdNone { - p.handleThreshold(dev, body) - } - - 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) - } - - return nil -} - -func (p *BatteryPlugin) handleThreshold(dev device.Sender, body BatteryBody) { - var message, urgency string - - switch body.ThresholdEvent { - case thresholdLow: - if !p.cfg.NotifyLow { - break - } - message = p.cfg.LowMessage - urgency = p.cfg.LowUrgency - case thresholdFull: - if !p.cfg.NotifyFull { - break - } - message = p.cfg.FullMessage - urgency = p.cfg.FullUrgency - default: - return - } - - if message != "" { - // Desktop notification. - plugin.RunCommandAsync(p.logger, "notify-send", - "-a", "KDE Connect", - "-u", urgency, - "-i", "battery", - dev.Name(), message, - ) - } - - // Emit event so watch / scripts can react. - if p.bus != nil { - p.bus.Publish(events.TypeBatteryThreshold, dev.ID(), map[string]any{ - "charge": body.CurrentCharge, - "charging": body.IsCharging, - "event": body.ThresholdEvent, - }) - } -} - -// 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{ - "request": true, - }) - dev.Send(pkt) - - // Send our local battery to the phone. - charge, charging, err := readLocalBattery() - if err != nil { - p.logger.Debug("local battery unavailable on connect", zap.Error(err)) - return - } - pkt, _ = protocol.NewPacket("kdeconnect.battery", BatteryBody{ - CurrentCharge: charge, - IsCharging: charging, - }) - dev.Send(pkt) -} - -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/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) + } +} diff --git a/internal/plugins/battery/handle.go b/internal/plugins/battery/handle.go new file mode 100644 index 0000000..39c73ab --- /dev/null +++ b/internal/plugins/battery/handle.go @@ -0,0 +1,125 @@ +package battery + +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 processes incoming battery packets. +func (p *BatteryPlugin) Handle(ctx context.Context, dev device.Sender, pkt *protocol.Packet) error { + switch pkt.Type { + case "kdeconnect.battery": + var body BatteryBody + if err := json.Unmarshal(pkt.Body, &body); err != nil { + 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 { + p.handleThreshold(dev, body) + } + + case "kdeconnect.battery.request": + // 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 + + switch body.ThresholdEvent { + case thresholdLow: + if !p.cfg.NotifyLow { + break + } + message = p.cfg.LowMessage + urgency = p.cfg.LowUrgency + case thresholdFull: + if !p.cfg.NotifyFull { + break + } + message = p.cfg.FullMessage + urgency = p.cfg.FullUrgency + default: + return + } + + if message != "" { + // Desktop notification. + plugin.RunCommandAsync(p.logger, "notify-send", + "-a", p.notifications.AppName(), + "-u", urgency, + "-i", "battery", + dev.Name(), message, + ) + } + + // Emit event so watch / scripts can react. + if p.bus != nil { + p.bus.Publish(events.TypeBatteryThreshold, dev.ID(), map[string]any{ + "charge": body.CurrentCharge, + "charging": body.IsCharging, + "event": body.ThresholdEvent, + }) + } +} + +// OnConnect requests the phone's battery and sends our local battery state. +func (p *BatteryPlugin) OnConnect(dev device.Sender) { + // 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) + + // Send our local battery to the phone. + charge, charging, err := readLocalBattery() + if err != nil { + p.logger.Debug("local battery unavailable on connect", zap.Error(err)) + return + } + pkt, _ = protocol.NewPacket("kdeconnect.battery", BatteryBody{ + CurrentCharge: charge, + IsCharging: charging, + }) + dev.Send(pkt) +} + +func (p *BatteryPlugin) OnDisconnect(_ device.Sender) {} 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 e779e0b..0000000 --- a/internal/plugins/clipboard/clipboard.go +++ /dev/null @@ -1,518 +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" - "go.uber.org/zap" -) - -type clipboardBackend int - -const ( - backendUnknown clipboardBackend = iota - backendWayland - backendX11 -) - -// ClipboardPlugin handles clipboard sync both directions. -type ClipboardPlugin struct { - 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) *ClipboardPlugin { - if logger == nil { - logger = zap.NewNop() - } - return &ClipboardPlugin{ - 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); 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) 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) - if err != nil { - return fmt.Errorf("clipboard: dial %s: %w", addr, 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) - } - defer f.Close() - - _, err = io.Copy(f, io.LimitReader(conn, size)) - if err != nil { - 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 6c7d5a7..0000000 --- a/internal/plugins/contacts/contacts.go +++ /dev/null @@ -1,591 +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) *ContactsPlugin { - dataHome := os.Getenv("XDG_DATA_HOME") - if dataHome == "" { - home, _ := os.UserHomeDir() - dataHome = filepath.Join(home, ".local", "share") - } - baseDir := filepath.Join(dataHome, "kcd", "contacts") - _ = 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/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/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 40438cb..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" ) @@ -63,7 +55,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 +78,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). @@ -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 a510634..0000000 --- a/internal/plugins/notification/notification.go +++ /dev/null @@ -1,449 +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" - "go.uber.org/zap" -) - -// NotificationPlugin handles incoming notifications and displays them on the desktop. -type NotificationPlugin struct { - 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) *NotificationPlugin { - p := &NotificationPlugin{ - 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"} -} - -// 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 "" - } - - 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) - 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 "" - } - 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) -} - -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/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 dee2cff..0000000 --- a/internal/plugins/pair/pair.go +++ /dev/null @@ -1,341 +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), - } -} - -// 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) - 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) - 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) - if len(vKey) > 16 { - vKey = vKey[:16] - } - 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) - 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) - } - - pkt, err := protocol.NewPairPacket(protocol.PairAccept) - if err != nil { - return err - } - - // Store our timestamp - p.mu.Lock() - p.pairingTimestamp[dev.ID()] = time.Now().Unix() - 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) - if len(vKey) > 16 { - vKey = vKey[:16] - } - 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) - 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) - 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/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/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 9b793a2..0000000 --- a/internal/plugins/share/share.go +++ /dev/null @@ -1,331 +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" - "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 { - 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) *SharePlugin { - return &SharePlugin{ - 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) - 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) - - 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/transfer.go b/internal/plugins/share/transfer.go index 9f35f75..b028311 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) @@ -82,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) } @@ -109,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() @@ -215,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/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 818c7ad..0000000 --- a/internal/plugins/sms/sms.go +++ /dev/null @@ -1,385 +0,0 @@ -package sms - -import ( - "context" - "crypto/tls" - "encoding/json" - "fmt" - "io" - "net" - "os" - "path/filepath" - "strconv" - "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" - "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 { - 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") - _ = os.MkdirAll(cacheDir, 0700) - - return &SMSPlugin{ - 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 bool `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": 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", "KDE Connect", - "-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) - } - 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) - 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) - } - defer f.Close() - - _, err = io.Copy(f, io.LimitReader(conn, size)) - if err != nil { - return fmt.Errorf("sms: receive attachment data: %w", err) - } - - return nil -} - -// --- SMS sending ----------------------------------------------------------- - -func (p *SMSPlugin) SendSMS(dev device.Sender, phoneNumber, message string) error { - body := map[string]any{ - "sendSms": true, - "phoneNumber": phoneNumber, - "messageBody": message, - } - 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/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/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/internal/plugins/telephony/telephony.go b/internal/plugins/telephony/telephony.go index 7fcad4f..5a6fa93 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,22 +14,28 @@ 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")), } } 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" } @@ -44,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 } @@ -69,14 +81,14 @@ 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: 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/plugins/telephony/telephony_test.go b/internal/plugins/telephony/telephony_test.go new file mode 100644 index 0000000..24fc42a --- /dev/null +++ b/internal/plugins/telephony/telephony_test.go @@ -0,0 +1,83 @@ +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" + "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) + // 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() + + 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) + } +} 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) + } + } +} 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) + } +} 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 new file mode 100644 index 0000000..52f330f --- /dev/null +++ b/internal/transport/sidechannel.go @@ -0,0 +1,73 @@ +package transport + +import ( + "context" + "crypto/tls" + "fmt" + "net" + "strconv" + "time" + + "github.com/bethropolis/kcd/internal/cert" + "go.uber.org/zap" +) + +// 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 + 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 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") + } + 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 WithIdleTimeout(conn, opts.IdleTimeout), nil +} diff --git a/internal/transport/sidechannel_test.go b/internal/transport/sidechannel_test.go new file mode 100644 index 0000000..3652f1e --- /dev/null +++ b/internal/transport/sidechannel_test.go @@ -0,0 +1,222 @@ +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.ListenConfig{}).Listen(context.Background(), "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.ListenConfig{}).Listen(context.Background(), "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) +} + +// 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) + } +} diff --git a/packaging/kcd.example.toml b/packaging/kcd.example.toml index dfff97c..1a22b5d 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,38 @@ 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 +# transfer_idle_timeout = "60s" # abort transfers silent longer than this + +[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 +220,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 +275,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 +284,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..d1381d0 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -9,16 +9,16 @@ 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. 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. @@ -75,454 +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) { - // Use a longer timeout for the listen operation - savedTimeout := c.Timeout - c.Timeout = 70 * time.Second - defer func() { c.Timeout = savedTimeout }() - - resp, err := c.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 -} - -// 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() + } + } +}