316 lines
8.8 KiB
Go
316 lines
8.8 KiB
Go
package main
|
|
|
|
import (
|
|
"encoding/json"
|
|
"flag"
|
|
"log"
|
|
"net/http"
|
|
"os"
|
|
"os/signal"
|
|
"syscall"
|
|
"time"
|
|
|
|
"github.com/fmartingr/ccrm/server/internal/api"
|
|
"github.com/fmartingr/ccrm/server/internal/config"
|
|
"github.com/fmartingr/ccrm/server/internal/db"
|
|
"github.com/fmartingr/ccrm/server/internal/ws"
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
func main() {
|
|
configPath := flag.String("config", "", "Path to config file")
|
|
flag.Parse()
|
|
|
|
cfg, err := config.Load(*configPath)
|
|
if err != nil {
|
|
log.Fatalf("Failed to load config: %v", err)
|
|
}
|
|
|
|
// Open database
|
|
database, err := db.Open(cfg.Database.Path)
|
|
if err != nil {
|
|
log.Fatalf("Failed to open database: %v", err)
|
|
}
|
|
defer database.Close()
|
|
|
|
// Ensure default machine exists
|
|
ensureDefaultMachine(database, cfg.Auth.DaemonAPIKey)
|
|
|
|
// Create WebSocket hub
|
|
hub := ws.NewHub()
|
|
|
|
// Set up message handlers
|
|
hub.SetDaemonMessageHandler(func(machineID string, msgType string, raw json.RawMessage) {
|
|
handleDaemonMessage(database, hub, machineID, msgType, raw)
|
|
})
|
|
|
|
hub.SetDashboardMessageHandler(func(msgType string, raw json.RawMessage) {
|
|
handleDashboardMessage(hub, msgType, raw)
|
|
})
|
|
|
|
go hub.Run()
|
|
|
|
// Create router
|
|
router := api.NewRouter(database, hub, cfg.Auth.DaemonAPIKey, cfg.Auth.DashboardToken)
|
|
|
|
// Serve static dashboard files if available
|
|
dashboardDir := "./dashboard/dist"
|
|
if _, err := os.Stat(dashboardDir); err == nil {
|
|
fs := http.FileServer(http.Dir(dashboardDir))
|
|
mux := http.NewServeMux()
|
|
mux.Handle("/api/", router)
|
|
mux.Handle("/ws/", router)
|
|
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
|
|
// Try serving the file, fall back to index.html for SPA routing
|
|
path := dashboardDir + r.URL.Path
|
|
if _, err := os.Stat(path); os.IsNotExist(err) && r.URL.Path != "/" {
|
|
http.ServeFile(w, r, dashboardDir+"/index.html")
|
|
return
|
|
}
|
|
fs.ServeHTTP(w, r)
|
|
})
|
|
router = mux
|
|
}
|
|
|
|
server := &http.Server{
|
|
Addr: cfg.Server.Listen,
|
|
Handler: router,
|
|
}
|
|
|
|
// Graceful shutdown
|
|
go func() {
|
|
sigCh := make(chan os.Signal, 1)
|
|
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
|
|
<-sigCh
|
|
log.Println("Shutting down server...")
|
|
server.Close()
|
|
}()
|
|
|
|
log.Printf("Server starting on %s", cfg.Server.Listen)
|
|
if err := server.ListenAndServe(); err != http.ErrServerClosed {
|
|
log.Fatalf("Server error: %v", err)
|
|
}
|
|
}
|
|
|
|
func ensureDefaultMachine(database *db.DB, apiKey string) {
|
|
var count int
|
|
err := database.QueryRow("SELECT COUNT(*) FROM machines WHERE api_key = ?", apiKey).Scan(&count)
|
|
if err != nil || count == 0 {
|
|
id := uuid.New().String()
|
|
_, err := database.Exec(
|
|
"INSERT OR IGNORE INTO machines (id, name, api_key, created_at) VALUES (?, ?, ?, ?)",
|
|
id, "default", apiKey, time.Now().UTC(),
|
|
)
|
|
if err != nil {
|
|
log.Printf("Warning: failed to create default machine: %v", err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func handleDaemonMessage(database *db.DB, hub *ws.Hub, machineID string, msgType string, raw json.RawMessage) {
|
|
// Update machine last_seen
|
|
database.Exec("UPDATE machines SET last_seen = ? WHERE name = ?", time.Now().UTC(), machineID)
|
|
|
|
switch msgType {
|
|
case "session.started":
|
|
var msg ws.SessionStartedMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
log.Printf("Failed to parse session.started: %v", err)
|
|
return
|
|
}
|
|
|
|
// Parse session data
|
|
var sessionData struct {
|
|
Name string `json:"name"`
|
|
TmuxSession string `json:"tmux_session"`
|
|
ProjectPath string `json:"project_path"`
|
|
ClaudeSessionID string `json:"claude_session_id"`
|
|
Status string `json:"status"`
|
|
IsWorktree bool `json:"is_worktree"`
|
|
WorktreeBranch string `json:"worktree_branch"`
|
|
CreatedAt string `json:"created_at"`
|
|
}
|
|
if err := json.Unmarshal(msg.Session, &sessionData); err != nil {
|
|
log.Printf("Failed to parse session data: %v", err)
|
|
return
|
|
}
|
|
|
|
// Get machine ID from DB
|
|
var dbMachineID string
|
|
err := database.QueryRow("SELECT id FROM machines WHERE name = ?", machineID).Scan(&dbMachineID)
|
|
if err != nil {
|
|
// Try default machine
|
|
err = database.QueryRow("SELECT id FROM machines WHERE name = 'default'").Scan(&dbMachineID)
|
|
if err != nil {
|
|
log.Printf("No machine found for %s: %v", machineID, err)
|
|
return
|
|
}
|
|
}
|
|
|
|
createdAt, _ := time.Parse(time.RFC3339, sessionData.CreatedAt)
|
|
if createdAt.IsZero() {
|
|
createdAt = time.Now().UTC()
|
|
}
|
|
|
|
s := &db.Session{
|
|
ID: uuid.New().String(),
|
|
MachineID: dbMachineID,
|
|
Name: sessionData.Name,
|
|
TmuxSession: sessionData.TmuxSession,
|
|
ProjectPath: sessionData.ProjectPath,
|
|
ClaudeSessionID: sessionData.ClaudeSessionID,
|
|
Status: sessionData.Status,
|
|
IsWorktree: sessionData.IsWorktree,
|
|
WorktreeBranch: sessionData.WorktreeBranch,
|
|
CreatedAt: createdAt,
|
|
}
|
|
|
|
if err := database.UpsertSession(s); err != nil {
|
|
log.Printf("Failed to upsert session: %v", err)
|
|
return
|
|
}
|
|
|
|
// Read back the actual session from DB (upsert may have kept old ID)
|
|
dbSession, err := database.GetSessionByName(s.Name, dbMachineID)
|
|
if err != nil {
|
|
log.Printf("Failed to read back session: %v", err)
|
|
dbSession = s
|
|
}
|
|
|
|
// Broadcast to dashboards
|
|
hub.BroadcastToDashboards(&ws.SessionUpdateMsg{
|
|
Type: "session.update",
|
|
Session: dbSession,
|
|
})
|
|
|
|
case "session.ended":
|
|
var msg ws.SessionEndedMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
return
|
|
}
|
|
|
|
var dbMachineID string
|
|
database.QueryRow("SELECT id FROM machines WHERE name = ? OR name = 'default'", machineID).Scan(&dbMachineID)
|
|
database.UpdateSessionStatus(msg.SessionName, dbMachineID, "stopped")
|
|
|
|
hub.BroadcastToDashboards(msg)
|
|
|
|
case "hook.event":
|
|
var msg ws.HookEventMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
log.Printf("Failed to parse hook.event: %v", err)
|
|
return
|
|
}
|
|
|
|
log.Printf("Hook event from daemon: session=%s", msg.SessionName)
|
|
|
|
// Find session in DB
|
|
var dbMachineID string
|
|
database.QueryRow("SELECT id FROM machines WHERE name = ? OR name = 'default'", machineID).Scan(&dbMachineID)
|
|
|
|
session, err := database.GetSessionByName(msg.SessionName, dbMachineID)
|
|
if err != nil {
|
|
log.Printf("Session %s not found in DB: %v", msg.SessionName, err)
|
|
return
|
|
}
|
|
|
|
// Parse event data for storage
|
|
var eventData struct {
|
|
EventName string `json:"hook_event_name"`
|
|
ToolName string `json:"tool_name"`
|
|
}
|
|
json.Unmarshal(msg.Event, &eventData)
|
|
|
|
log.Printf("Storing event: session_id=%s event=%s tool=%s", session.ID, eventData.EventName, eventData.ToolName)
|
|
|
|
// Store event
|
|
database.InsertEvent(session.ID, eventData.EventName, eventData.ToolName, string(msg.Event))
|
|
|
|
// Update session status based on event
|
|
var newStatus string
|
|
switch eventData.EventName {
|
|
case "SessionStart":
|
|
newStatus = "active"
|
|
case "SessionEnd":
|
|
newStatus = "stopped"
|
|
case "UserPromptSubmit":
|
|
newStatus = "active"
|
|
case "Stop":
|
|
newStatus = "idle"
|
|
case "PermissionRequest":
|
|
newStatus = "waiting_permission"
|
|
case "Notification":
|
|
newStatus = "idle"
|
|
}
|
|
if newStatus != "" {
|
|
database.UpdateSessionStatus(msg.SessionName, dbMachineID, newStatus)
|
|
session.Status = newStatus
|
|
// Broadcast session status update
|
|
hub.BroadcastToDashboards(&ws.SessionUpdateMsg{
|
|
Type: "session.update",
|
|
Session: session,
|
|
})
|
|
}
|
|
|
|
// Broadcast new event to dashboards
|
|
hub.BroadcastToDashboards(&ws.EventNewMsg{
|
|
Type: "event.new",
|
|
SessionName: msg.SessionName,
|
|
Event: json.RawMessage(msg.Event),
|
|
})
|
|
|
|
case "prompt.ack":
|
|
// Forward acknowledgment to dashboards
|
|
hub.BroadcastToDashboards(json.RawMessage(raw))
|
|
|
|
case "heartbeat":
|
|
var msg ws.HeartbeatMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
log.Printf("Failed to parse heartbeat: %v", err)
|
|
return
|
|
}
|
|
|
|
// Reconcile: mark any DB sessions not in the live list as stopped
|
|
var dbMachineID string
|
|
err := database.QueryRow("SELECT id FROM machines WHERE name = ? OR name = 'default'", machineID).Scan(&dbMachineID)
|
|
if err == nil {
|
|
if err := database.ReconcileSessions(dbMachineID, msg.Sessions); err != nil {
|
|
log.Printf("Failed to reconcile sessions: %v", err)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func handleDashboardMessage(hub *ws.Hub, msgType string, raw json.RawMessage) {
|
|
switch msgType {
|
|
case "prompt.send":
|
|
var msg ws.DashboardPromptSendMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
return
|
|
}
|
|
|
|
// Forward to daemon
|
|
daemonIDs := hub.GetConnectedDaemonIDs()
|
|
if len(daemonIDs) > 0 {
|
|
hub.SendToDaemon(daemonIDs[0], map[string]string{
|
|
"type": "prompt.send",
|
|
"session_name": msg.SessionName,
|
|
"prompt": msg.Prompt,
|
|
})
|
|
}
|
|
|
|
case "session.kill":
|
|
var msg ws.DashboardSessionKillMsg
|
|
if err := json.Unmarshal(raw, &msg); err != nil {
|
|
return
|
|
}
|
|
|
|
daemonIDs := hub.GetConnectedDaemonIDs()
|
|
if len(daemonIDs) > 0 {
|
|
hub.SendToDaemon(daemonIDs[0], map[string]string{
|
|
"type": "session.kill",
|
|
"session_name": msg.SessionName,
|
|
})
|
|
}
|
|
}
|
|
}
|