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()