From 4246ac532191aa541bd2a12791dc826fd0042e46 Mon Sep 17 00:00:00 2001 From: KoheiYamashita Date: Sat, 21 Feb 2026 16:22:11 +0900 Subject: [PATCH] feat: add client type tracking and broadcast fallback for disconnected WS clients When the main Android app's WebSocket connection is down, messages are now delivered via Android broadcast intent (am broadcast) so the app can save them to DB and show notifications. WebSocket sessions now carry a client_type parameter ("main" for the Android app) to distinguish between different client types for future Google Assistant replacement clients. Heartbeat always targets the last "main" session. - Add LastMainChannel to state for heartbeat targeting - Add clientTypes map to WebSocketChannel (retained after disconnect) - Add broadcast fallback in WebSocket Send for "main" clients - Add pkg/broadcast package using am broadcast IPC - Add AgentMessageReceiver + NotificationHelper on Android side - Request POST_NOTIFICATIONS runtime permission on Android 13+ - Share state.Manager between AgentLoop and HeartbeatService Co-Authored-By: Claude Opus 4.6 --- android/app/src/main/AndroidManifest.xml | 9 ++ .../java/io/picoclaw/android/MainActivity.kt | 21 +++++ .../java/io/picoclaw/android/PicoClawApp.kt | 2 + .../android/receiver/AgentMessageReceiver.kt | 62 ++++++++++++++ .../android/receiver/NotificationHelper.kt | 47 +++++++++++ .../core/data/remote/WebSocketClient.kt | 2 +- cmd/picoclaw/main.go | 1 + pkg/agent/loop.go | 42 +++++++--- pkg/broadcast/broadcast.go | 51 +++++++++++ pkg/channels/websocket.go | 84 +++++++++++++++---- pkg/heartbeat/service.go | 27 ++++-- pkg/heartbeat/service_test.go | 16 ++-- pkg/state/state.go | 32 +++++++ 13 files changed, 352 insertions(+), 44 deletions(-) create mode 100644 android/app/src/main/java/io/picoclaw/android/receiver/AgentMessageReceiver.kt create mode 100644 android/app/src/main/java/io/picoclaw/android/receiver/NotificationHelper.kt create mode 100644 pkg/broadcast/broadcast.go diff --git a/android/app/src/main/AndroidManifest.xml b/android/app/src/main/AndroidManifest.xml index 0399f91cc..f170e9b1e 100644 --- a/android/app/src/main/AndroidManifest.xml +++ b/android/app/src/main/AndroidManifest.xml @@ -4,6 +4,7 @@ + @@ -38,6 +39,14 @@ + + + + + + diff --git a/android/app/src/main/java/io/picoclaw/android/MainActivity.kt b/android/app/src/main/java/io/picoclaw/android/MainActivity.kt index 6d8fa5241..babec4acd 100644 --- a/android/app/src/main/java/io/picoclaw/android/MainActivity.kt +++ b/android/app/src/main/java/io/picoclaw/android/MainActivity.kt @@ -1,9 +1,14 @@ package io.picoclaw.android +import android.Manifest +import android.content.pm.PackageManager +import android.os.Build import android.os.Bundle import androidx.activity.ComponentActivity import androidx.activity.compose.setContent import androidx.activity.enableEdgeToEdge +import androidx.activity.result.contract.ActivityResultContracts +import androidx.core.content.ContextCompat import androidx.navigation.compose.NavHost import androidx.navigation.compose.composable import androidx.navigation.compose.rememberNavController @@ -13,8 +18,14 @@ import io.picoclaw.android.feature.chat.screen.SettingsScreen import io.picoclaw.android.navigation.NavRoutes class MainActivity : ComponentActivity() { + + private val notificationPermissionLauncher = registerForActivityResult( + ActivityResultContracts.RequestPermission() + ) { /* User choice recorded; no further action needed. */ } + override fun onCreate(savedInstanceState: Bundle?) { super.onCreate(savedInstanceState) + requestNotificationPermissionIfNeeded() enableEdgeToEdge() setContent { PicoClawTheme { @@ -34,4 +45,14 @@ class MainActivity : ComponentActivity() { } } } + + private fun requestNotificationPermissionIfNeeded() { + if (Build.VERSION.SDK_INT >= Build.VERSION_CODES.TIRAMISU) { + if (ContextCompat.checkSelfPermission(this, Manifest.permission.POST_NOTIFICATIONS) + != PackageManager.PERMISSION_GRANTED + ) { + notificationPermissionLauncher.launch(Manifest.permission.POST_NOTIFICATIONS) + } + } + } } diff --git a/android/app/src/main/java/io/picoclaw/android/PicoClawApp.kt b/android/app/src/main/java/io/picoclaw/android/PicoClawApp.kt index b4d3be91e..410f02d2c 100644 --- a/android/app/src/main/java/io/picoclaw/android/PicoClawApp.kt +++ b/android/app/src/main/java/io/picoclaw/android/PicoClawApp.kt @@ -2,6 +2,7 @@ package io.picoclaw.android import android.app.Application import io.picoclaw.android.di.appModule +import io.picoclaw.android.receiver.NotificationHelper import org.koin.android.ext.koin.androidContext import org.koin.core.context.startKoin @@ -12,5 +13,6 @@ class PicoClawApp : Application() { androidContext(this@PicoClawApp) modules(appModule) } + NotificationHelper.createNotificationChannel(this) } } diff --git a/android/app/src/main/java/io/picoclaw/android/receiver/AgentMessageReceiver.kt b/android/app/src/main/java/io/picoclaw/android/receiver/AgentMessageReceiver.kt new file mode 100644 index 000000000..9da27fc50 --- /dev/null +++ b/android/app/src/main/java/io/picoclaw/android/receiver/AgentMessageReceiver.kt @@ -0,0 +1,62 @@ +package io.picoclaw.android.receiver + +import android.content.BroadcastReceiver +import android.content.Context +import android.content.Intent +import android.util.Log +import io.picoclaw.android.core.data.local.dao.MessageDao +import io.picoclaw.android.core.data.mapper.MessageMapper +import io.picoclaw.android.core.data.remote.dto.WsOutgoing +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.launch +import kotlinx.serialization.json.Json +import org.koin.core.component.KoinComponent +import org.koin.core.component.inject + +class AgentMessageReceiver : BroadcastReceiver(), KoinComponent { + + private val messageDao: MessageDao by inject() + private val scope: CoroutineScope by inject() + private val json = Json { ignoreUnknownKeys = true } + + override fun onReceive(context: Context, intent: Intent) { + val messageJson = intent.getStringExtra("message") ?: return + + Log.d(TAG, "Received broadcast message: ${messageJson.take(100)}") + + val pendingResult = goAsync() + + try { + val msg = json.decodeFromString(messageJson) + + // Skip ephemeral status messages + if (msg.type == "status" || msg.type == "status_end") { + pendingResult.finish() + return + } + + // Save to DB then finish + val entity = MessageMapper.toEntity(msg) + scope.launch { + try { + messageDao.insert(entity) + } catch (e: Exception) { + Log.e(TAG, "Failed to insert message to DB", e) + } finally { + pendingResult.finish() + } + } + + // Show notification (synchronous, safe to call here) + NotificationHelper.showMessageNotification(context, msg.content) + } catch (e: Exception) { + Log.e(TAG, "Failed to process broadcast message", e) + pendingResult.finish() + } + } + + companion object { + private const val TAG = "AgentMessageReceiver" + const val ACTION = "io.picoclaw.android.AGENT_MESSAGE" + } +} diff --git a/android/app/src/main/java/io/picoclaw/android/receiver/NotificationHelper.kt b/android/app/src/main/java/io/picoclaw/android/receiver/NotificationHelper.kt new file mode 100644 index 000000000..9d0501507 --- /dev/null +++ b/android/app/src/main/java/io/picoclaw/android/receiver/NotificationHelper.kt @@ -0,0 +1,47 @@ +package io.picoclaw.android.receiver + +import android.app.NotificationChannel +import android.app.NotificationManager +import android.app.PendingIntent +import android.content.Context +import androidx.core.app.NotificationCompat + +object NotificationHelper { + + private const val CHANNEL_ID = "picoclaw_messages" + private const val CHANNEL_NAME = "Agent Messages" + private const val NOTIFICATION_ID = 1001 + + fun createNotificationChannel(context: Context) { + val channel = NotificationChannel( + CHANNEL_ID, + CHANNEL_NAME, + NotificationManager.IMPORTANCE_DEFAULT + ).apply { + description = "Messages from PicoClaw agent" + } + val manager = context.getSystemService(NotificationManager::class.java) + manager.createNotificationChannel(channel) + } + + fun showMessageNotification(context: Context, content: String) { + val launchIntent = context.packageManager.getLaunchIntentForPackage(context.packageName) + val pendingIntent = PendingIntent.getActivity( + context, 0, launchIntent, + PendingIntent.FLAG_UPDATE_CURRENT or PendingIntent.FLAG_IMMUTABLE + ) + + val notification = NotificationCompat.Builder(context, CHANNEL_ID) + .setSmallIcon(android.R.drawable.ic_dialog_info) + .setContentTitle("PicoClaw") + .setContentText(content.take(200)) + .setStyle(NotificationCompat.BigTextStyle().bigText(content.take(1000))) + .setPriority(NotificationCompat.PRIORITY_DEFAULT) + .setContentIntent(pendingIntent) + .setAutoCancel(true) + .build() + + val manager = context.getSystemService(NotificationManager::class.java) + manager.notify(NOTIFICATION_ID, notification) + } +} diff --git a/android/core/data/src/main/java/io/picoclaw/android/core/data/remote/WebSocketClient.kt b/android/core/data/src/main/java/io/picoclaw/android/core/data/remote/WebSocketClient.kt index 7df89a654..b560ad45c 100644 --- a/android/core/data/src/main/java/io/picoclaw/android/core/data/remote/WebSocketClient.kt +++ b/android/core/data/src/main/java/io/picoclaw/android/core/data/remote/WebSocketClient.kt @@ -48,7 +48,7 @@ class WebSocketClient( while (isActive) { try { _connectionState.value = ConnectionState.CONNECTING - val url = "$wsUrl?client_id=$clientId" + val url = "$wsUrl?client_id=$clientId&client_type=main" client.webSocket(url) { session = this _connectionState.value = ConnectionState.CONNECTED diff --git a/cmd/picoclaw/main.go b/cmd/picoclaw/main.go index d84cec019..a48cdc49d 100644 --- a/cmd/picoclaw/main.go +++ b/cmd/picoclaw/main.go @@ -566,6 +566,7 @@ func gatewayCmd() { cfg.DataPath(), cfg.Heartbeat.Interval, cfg.Heartbeat.Enabled, + agentLoop.StateManager(), ) heartbeatService.SetBus(msgBus) heartbeatService.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 350d4264c..3ab643c87 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -58,16 +58,17 @@ type activeProcess struct { // processOptions configures how a message is processed type processOptions struct { - SessionKey string // Session identifier for history/context - Channel string // Target channel for tool execution - ChatID string // Target chat ID for tool execution - UserMessage string // User message content (may include prefix) - Media []string // Base64 data URLs for images - DefaultResponse string // Response when LLM returns empty - EnableSummary bool // Whether to trigger summarization - SendResponse bool // Whether to send response via bus - NoHistory bool // If true, don't load session history (for heartbeat) - InputMode string // "voice" or "text" + SessionKey string // Session identifier for history/context + Channel string // Target channel for tool execution + ChatID string // Target chat ID for tool execution + UserMessage string // User message content (may include prefix) + Media []string // Base64 data URLs for images + DefaultResponse string // Response when LLM returns empty + EnableSummary bool // Whether to trigger summarization + SendResponse bool // Whether to send response via bus + NoHistory bool // If true, don't load session history (for heartbeat) + InputMode string // "voice" or "text" + Metadata map[string]string // Channel metadata (e.g. client_type) } // createToolRegistry creates a tool registry with common tools. @@ -296,6 +297,12 @@ func (al *AgentLoop) SetChannelManager(cm *channels.Manager) { al.channelManager = cm } +// StateManager returns the state manager used by this agent loop. +// This allows sharing the same instance with other services (e.g. heartbeat). +func (al *AgentLoop) StateManager() *state.Manager { + return al.state +} + // RecordLastChannel records the last active channel for this workspace. // This uses the atomic state save mechanism to prevent data loss on crash. func (al *AgentLoop) RecordLastChannel(channel string) error { @@ -394,6 +401,7 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage) EnableSummary: true, SendResponse: false, InputMode: inputMode, + Metadata: msg.Metadata, }) } @@ -457,8 +465,18 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str // Don't record internal channels (cli, system, subagent) if !constants.IsInternalChannel(opts.Channel) { channelKey := fmt.Sprintf("%s:%s", opts.Channel, opts.ChatID) - if err := al.RecordLastChannel(channelKey); err != nil { - logger.WarnCF("agent", "Failed to record last channel: %v", map[string]interface{}{"error": err.Error()}) + clientType := "" + if opts.Metadata != nil { + clientType = opts.Metadata["client_type"] + } + if clientType != "" { + if err := al.state.SetLastChannelWithType(channelKey, clientType); err != nil { + logger.WarnCF("agent", "Failed to record last channel: %v", map[string]interface{}{"error": err.Error()}) + } + } else { + if err := al.RecordLastChannel(channelKey); err != nil { + logger.WarnCF("agent", "Failed to record last channel: %v", map[string]interface{}{"error": err.Error()}) + } } } } diff --git a/pkg/broadcast/broadcast.go b/pkg/broadcast/broadcast.go new file mode 100644 index 000000000..fdcee1634 --- /dev/null +++ b/pkg/broadcast/broadcast.go @@ -0,0 +1,51 @@ +package broadcast + +import ( + "encoding/json" + "fmt" + "os/exec" + + "github.com/sipeed/picoclaw/pkg/logger" +) + +const ( + // Action is the intent action the Android app listens for. + Action = "io.picoclaw.android.AGENT_MESSAGE" + // Package is the Android app package name. + Package = "io.picoclaw.android" +) + +// Message represents a message to send via Android broadcast. +type Message struct { + Content string `json:"content"` + Type string `json:"type,omitempty"` +} + +// Send sends a message to the Android app via am broadcast. +// This works because the Go server runs inside Termux on the same device. +func Send(msg Message) error { + data, err := json.Marshal(msg) + if err != nil { + return fmt.Errorf("failed to marshal broadcast message: %w", err) + } + + cmd := exec.Command("am", "broadcast", + "-a", Action, + "-p", Package, + "--es", "message", string(data), + ) + + output, err := cmd.CombinedOutput() + if err != nil { + logger.ErrorCF("broadcast", "am broadcast failed", map[string]interface{}{ + "error": err.Error(), + "output": string(output), + }) + return fmt.Errorf("am broadcast failed: %w (%s)", err, string(output)) + } + + logger.InfoCF("broadcast", "Broadcast sent", map[string]interface{}{ + "content_len": len(msg.Content), + }) + return nil +} diff --git a/pkg/channels/websocket.go b/pkg/channels/websocket.go index 0c029d399..a9b988a54 100644 --- a/pkg/channels/websocket.go +++ b/pkg/channels/websocket.go @@ -9,6 +9,7 @@ import ( "github.com/google/uuid" "github.com/gorilla/websocket" + "github.com/sipeed/picoclaw/pkg/broadcast" "github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/logger" @@ -32,14 +33,15 @@ type wsOutgoing struct { // connections from clients (e.g. a Google Assistant replacement APK). type WebSocketChannel struct { *BaseChannel - config config.WebSocketConfig - server *http.Server - upgrader websocket.Upgrader - clients map[*websocket.Conn]string // conn → clientID - chatConns map[string]*websocket.Conn // chatID → conn - mu sync.RWMutex - ctx context.Context - cancel context.CancelFunc + config config.WebSocketConfig + server *http.Server + upgrader websocket.Upgrader + clients map[*websocket.Conn]string // conn → clientID + chatConns map[string]*websocket.Conn // chatID → conn + clientTypes map[string]string // chatID → clientType (retained after disconnect) + mu sync.RWMutex + ctx context.Context + cancel context.CancelFunc } func NewWebSocketChannel(cfg config.WebSocketConfig, msgBus *bus.MessageBus) (*WebSocketChannel, error) { @@ -51,8 +53,9 @@ func NewWebSocketChannel(cfg config.WebSocketConfig, msgBus *bus.MessageBus) (*W upgrader: websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, }, - clients: make(map[*websocket.Conn]string), - chatConns: make(map[string]*websocket.Conn), + clients: make(map[*websocket.Conn]string), + chatConns: make(map[string]*websocket.Conn), + clientTypes: make(map[string]string), }, nil } @@ -106,6 +109,7 @@ func (c *WebSocketChannel) Stop(ctx context.Context) error { } c.clients = make(map[*websocket.Conn]string) c.chatConns = make(map[string]*websocket.Conn) + c.clientTypes = make(map[string]string) c.mu.Unlock() if c.server != nil { @@ -128,10 +132,12 @@ func (c *WebSocketChannel) Send(ctx context.Context, msg bus.OutboundMessage) er c.mu.RLock() conn, ok := c.chatConns[msg.ChatID] + clientType := c.clientTypes[msg.ChatID] // Read under lock to avoid data race c.mu.RUnlock() if !ok { - return fmt.Errorf("no connection for chat %s", msg.ChatID) + // Connection not found — try broadcast fallback for "main" clients. + return c.maybeBroadcast(msg, clientType, fmt.Errorf("no connection for chat %s", msg.ChatID)) } out := wsOutgoing{Content: msg.Content, Type: msg.Type} @@ -145,7 +151,7 @@ func (c *WebSocketChannel) Send(ctx context.Context, msg bus.OutboundMessage) er // Verify connection still exists (may have been removed during cleanup). if _, exists := c.clients[conn]; !exists { - return fmt.Errorf("connection for chat %s no longer active", msg.ChatID) + return c.maybeBroadcast(msg, c.clientTypes[msg.ChatID], fmt.Errorf("connection for chat %s no longer active", msg.ChatID)) } if err := conn.WriteMessage(websocket.TextMessage, data); err != nil { @@ -153,12 +159,41 @@ func (c *WebSocketChannel) Send(ctx context.Context, msg bus.OutboundMessage) er "chat_id": msg.ChatID, "error": err.Error(), }) - return err + return c.maybeBroadcast(msg, c.clientTypes[msg.ChatID], err) } return nil } +// maybeBroadcast sends a message via Android broadcast if the disconnected +// client is of type "main". Status/status_end messages are ephemeral and +// skipped. Returns the original error if broadcast is not applicable. +// clientType must be read under lock by the caller to avoid data races. +func (c *WebSocketChannel) maybeBroadcast(msg bus.OutboundMessage, clientType string, originalErr error) error { + // Status messages are ephemeral — don't broadcast. + if msg.Type == "status" || msg.Type == "status_end" { + return originalErr + } + + // Only broadcast for "main" type clients. + if clientType != "main" { + return originalErr + } + + logger.InfoCF("websocket", "Using broadcast fallback for disconnected main client", map[string]interface{}{ + "chat_id": msg.ChatID, + "content_len": len(msg.Content), + }) + + if err := broadcast.Send(broadcast.Message{ + Content: msg.Content, + Type: msg.Type, + }); err != nil { + return fmt.Errorf("ws send failed and broadcast fallback also failed: %w", err) + } + return nil +} + func (c *WebSocketChannel) handleWS(w http.ResponseWriter, r *http.Request) { conn, err := c.upgrader.Upgrade(w, r, nil) if err != nil { @@ -173,8 +208,14 @@ func (c *WebSocketChannel) handleWS(w http.ResponseWriter, r *http.Request) { clientID = uuid.New().String() } + clientType := r.URL.Query().Get("client_type") + if clientType == "" { + clientType = "main" // Default for backward compatibility + } + logger.InfoCF("websocket", "New WebSocket connection", map[string]interface{}{ "client_id": clientID, + "client_type": clientType, "remote_addr": r.RemoteAddr, }) @@ -187,12 +228,22 @@ func (c *WebSocketChannel) handleWS(w http.ResponseWriter, r *http.Request) { } c.clients[conn] = clientID c.chatConns[chatID] = conn + c.clientTypes[chatID] = clientType c.mu.Unlock() - go c.readPump(conn, clientID, chatID) + go c.readPump(conn, clientID, chatID, clientType) } -func (c *WebSocketChannel) readPump(conn *websocket.Conn, clientID, chatID string) { +// GetClientType returns the client type for a given chatID. +// Returns empty string if unknown. The value is retained after disconnect +// so that the broadcast fallback can check the type. +func (c *WebSocketChannel) GetClientType(chatID string) string { + c.mu.RLock() + defer c.mu.RUnlock() + return c.clientTypes[chatID] +} + +func (c *WebSocketChannel) readPump(conn *websocket.Conn, clientID, chatID, clientType string) { defer func() { c.mu.Lock() delete(c.clients, conn) @@ -258,7 +309,8 @@ func (c *WebSocketChannel) readPump(conn *websocket.Conn, clientID, chatID strin inputMode = "text" } metadata := map[string]string{ - "input_mode": inputMode, + "input_mode": inputMode, + "client_type": clientType, } c.HandleMessage(senderID, chatID, content, media, metadata) } diff --git a/pkg/heartbeat/service.go b/pkg/heartbeat/service.go index 0cfc50684..9eb671e1e 100644 --- a/pkg/heartbeat/service.go +++ b/pkg/heartbeat/service.go @@ -44,8 +44,9 @@ type HeartbeatService struct { stopChan chan struct{} } -// NewHeartbeatService creates a new heartbeat service -func NewHeartbeatService(workspace string, dataDir string, intervalMinutes int, enabled bool) *HeartbeatService { +// NewHeartbeatService creates a new heartbeat service. +// If stateManager is nil, a new state.Manager is created internally (for tests). +func NewHeartbeatService(workspace string, dataDir string, intervalMinutes int, enabled bool, stateManager *state.Manager) *HeartbeatService { // Apply minimum interval if intervalMinutes < minIntervalMinutes && intervalMinutes != 0 { intervalMinutes = minIntervalMinutes @@ -55,12 +56,16 @@ func NewHeartbeatService(workspace string, dataDir string, intervalMinutes int, intervalMinutes = defaultIntervalMinutes } + if stateManager == nil { + stateManager = state.NewManager(dataDir) + } + return &HeartbeatService{ workspace: workspace, dataDir: dataDir, interval: time.Duration(intervalMinutes) * time.Minute, enabled: enabled, - state: state.NewManager(dataDir), + state: stateManager, } } @@ -172,8 +177,13 @@ func (hs *HeartbeatService) executeHeartbeat() { return } - // Get last channel info for context - lastChannel := hs.state.GetLastChannel() + // Get last channel info for context. + // Prefer the last "main" client session so heartbeat always targets the + // primary Android app, even if a non-main client was active more recently. + lastChannel := hs.state.GetLastMainChannel() + if lastChannel == "" { + lastChannel = hs.state.GetLastChannel() // Backward-compat fallback + } channel, chatID := hs.parseLastChannel(lastChannel) // Debug log for channel resolution @@ -295,8 +305,11 @@ func (hs *HeartbeatService) sendResponse(response string) { return } - // Get last channel from state - lastChannel := hs.state.GetLastChannel() + // Get last channel from state (prefer "main" client session) + lastChannel := hs.state.GetLastMainChannel() + if lastChannel == "" { + lastChannel = hs.state.GetLastChannel() // Backward-compat fallback + } if lastChannel == "" { hs.logInfo("No last channel recorded, heartbeat result not sent") return diff --git a/pkg/heartbeat/service_test.go b/pkg/heartbeat/service_test.go index b55f3e601..41fc68f59 100644 --- a/pkg/heartbeat/service_test.go +++ b/pkg/heartbeat/service_test.go @@ -16,7 +16,7 @@ func TestExecuteHeartbeat_Async(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) hs.stopChan = make(chan struct{}) // Enable for testing asyncCalled := false @@ -54,7 +54,7 @@ func TestExecuteHeartbeat_Error(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) hs.stopChan = make(chan struct{}) // Enable for testing hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { @@ -92,7 +92,7 @@ func TestExecuteHeartbeat_Silent(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) hs.stopChan = make(chan struct{}) // Enable for testing hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { @@ -130,7 +130,7 @@ func TestHeartbeatService_StartStop(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 1, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 1, true, nil) err = hs.Start() if err != nil { @@ -149,7 +149,7 @@ func TestHeartbeatService_Disabled(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 1, false) + hs := NewHeartbeatService(tmpDir, tmpDir, 1, false, nil) if hs.enabled != false { t.Error("Expected service to be disabled") @@ -166,7 +166,7 @@ func TestExecuteHeartbeat_NilResult(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) hs.stopChan = make(chan struct{}) // Enable for testing hs.SetHandler(func(prompt, channel, chatID string) *tools.ToolResult { @@ -188,7 +188,7 @@ func TestLogPath(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) // Write a log entry hs.log("INFO", "Test log entry") @@ -208,7 +208,7 @@ func TestHeartbeatFilePath(t *testing.T) { } defer os.RemoveAll(tmpDir) - hs := NewHeartbeatService(tmpDir, tmpDir, 30, true) + hs := NewHeartbeatService(tmpDir, tmpDir, 30, true, nil) // Trigger default template creation hs.buildPrompt() diff --git a/pkg/state/state.go b/pkg/state/state.go index 0bb9cd497..cc7473fbb 100644 --- a/pkg/state/state.go +++ b/pkg/state/state.go @@ -19,6 +19,10 @@ type State struct { // LastChatID is the last chat ID used for communication LastChatID string `json:"last_chat_id,omitempty"` + // LastMainChannel is the last channel used by a "main" client type. + // Used by heartbeat to always target the main (Android app) session. + LastMainChannel string `json:"last_main_channel,omitempty"` + // Timestamp is the last time this state was updated Timestamp time.Time `json:"timestamp"` } @@ -114,6 +118,34 @@ func (sm *Manager) GetLastChatID() string { return sm.state.LastChatID } +// SetLastChannelWithType atomically updates the last channel and, if the +// client type is "main", also updates LastMainChannel. Both fields are +// persisted in a single atomic write. +func (sm *Manager) SetLastChannelWithType(channel, clientType string) error { + sm.mu.Lock() + defer sm.mu.Unlock() + + sm.state.LastChannel = channel + sm.state.Timestamp = time.Now() + + if clientType == "main" { + sm.state.LastMainChannel = channel + } + + if err := sm.saveAtomic(); err != nil { + return fmt.Errorf("failed to save state atomically: %w", err) + } + + return nil +} + +// GetLastMainChannel returns the last channel used by a "main" client type. +func (sm *Manager) GetLastMainChannel() string { + sm.mu.RLock() + defer sm.mu.RUnlock() + return sm.state.LastMainChannel +} + // GetTimestamp returns the timestamp of the last state update. func (sm *Manager) GetTimestamp() time.Time { sm.mu.RLock()