package transport import ( "context" "fmt" "net" "strconv" "sync" toki "git.toki-labs.com/toki/common-proto-socket/go" "go.uber.org/zap" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/structpb" edgenode "iop/apps/edge/internal/node" iop "iop/proto/gen/iop" ) const ( heartbeatIntervalSec = 30 heartbeatWaitSec = 10 ) func edgeParserMap() toki.ParserMap { return toki.ParserMap{ toki.TypeNameOf(&iop.RunEvent{}): func(b []byte) (proto.Message, error) { m := &iop.RunEvent{} return m, proto.Unmarshal(b, m) }, toki.TypeNameOf(&iop.RegisterRequest{}): func(b []byte) (proto.Message, error) { m := &iop.RegisterRequest{} return m, proto.Unmarshal(b, m) }, } } // Server wraps proto-socket TcpServer and manages node connections. type Server struct { tcp *toki.TcpServer listen string registry *edgenode.Registry nodeStore *edgenode.NodeStore logger *zap.Logger handlerMu sync.RWMutex onRunEvent func(*iop.RunEvent) } func NewServer(listen string, registry *edgenode.Registry, nodeStore *edgenode.NodeStore, logger *zap.Logger) (*Server, error) { host, portStr, err := net.SplitHostPort(listen) if err != nil { return nil, err } port, err := strconv.Atoi(portStr) if err != nil { return nil, err } s := &Server{listen: listen, registry: registry, nodeStore: nodeStore, logger: logger} s.tcp = toki.NewTcpServer(host, port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, heartbeatIntervalSec, heartbeatWaitSec, edgeParserMap()) }) s.tcp.OnClientConnected = s.onNodeConnected return s, nil } func (s *Server) Start(ctx context.Context) error { if err := s.tcp.Start(ctx); err != nil { return err } s.logger.Info("edge listening for nodes", zap.String("addr", s.listen)) return nil } func (s *Server) Stop() error { return s.tcp.Stop() } func (s *Server) SetRunEventHandler(handler func(*iop.RunEvent)) { s.handlerMu.Lock() s.onRunEvent = handler s.handlerMu.Unlock() } func (s *Server) onNodeConnected(client *toki.TcpClient) { s.logger.Info("node connection established") toki.AddListenerTyped[*iop.RunEvent](&client.Communicator, func(e *iop.RunEvent) { s.logger.Debug("run event received", zap.String("run_id", e.GetRunId()), zap.String("type", e.GetType()), ) s.handlerMu.RLock() handler := s.onRunEvent s.handlerMu.RUnlock() if handler != nil { handler(e) } }) toki.AddRequestListenerTyped[*iop.RegisterRequest, *iop.RegisterResponse]( &client.Communicator, func(req *iop.RegisterRequest) (*iop.RegisterResponse, error) { rec, ok := s.nodeStore.FindByToken(req.GetToken()) if !ok { s.logger.Warn("unknown token", zap.String("token_prefix", safePrefix(req.GetToken()))) return &iop.RegisterResponse{Accepted: false, Reason: "unknown token"}, nil } if _, ok := s.registry.Get(rec.ID); ok { s.logger.Warn("node already connected", zap.String("node_id", rec.ID)) return &iop.RegisterResponse{Accepted: false, Reason: "node already connected"}, nil } cfg, err := buildConfigPayload(rec) if err != nil { s.logger.Error("build config payload failed", zap.String("node_id", rec.ID), zap.Error(err), ) return &iop.RegisterResponse{Accepted: false, Reason: "internal config error"}, nil } entry := &edgenode.NodeEntry{ NodeID: rec.ID, Alias: rec.Alias, Client: client, } client.AddDisconnectListener(func(_ *toki.TcpClient) { s.registry.Unregister(rec.ID) s.logger.Info("node unregistered", zap.String("node_id", rec.ID)) }) s.registry.Register(entry) s.logger.Info("node registered", zap.String("node_id", rec.ID), zap.String("alias", rec.Alias), ) return &iop.RegisterResponse{ Accepted: true, NodeId: rec.ID, Alias: rec.Alias, Config: cfg, }, nil }, ) } func safePrefix(s string) string { if len(s) > 8 { return s[:8] + "..." } return s } func buildConfigPayload(rec *edgenode.NodeRecord) (*iop.NodeConfigPayload, error) { payload := &iop.NodeConfigPayload{ Runtime: &iop.NodeRuntimeConfig{ Concurrency: int32(rec.Runtime.Concurrency), WorkspaceRoot: rec.Runtime.WorkspaceRoot, }, } addAdapter := func(typ string, enabled bool, settings map[string]any) error { if !enabled { return nil } st, err := structpb.NewStruct(settings) if err != nil { return fmt.Errorf("buildConfigPayload: %s: %w", typ, err) } payload.Adapters = append(payload.Adapters, &iop.AdapterConfig{ Type: typ, Enabled: true, Settings: st, }) return nil } if err := addAdapter("mock", true, nil); err != nil { return nil, err } if err := addAdapter("ollama", rec.Adapters.Ollama.Enabled, map[string]any{ "base_url": rec.Adapters.Ollama.BaseURL, }); err != nil { return nil, err } if err := addAdapter("vllm", rec.Adapters.Vllm.Enabled, map[string]any{ "endpoint": rec.Adapters.Vllm.Endpoint, }); err != nil { return nil, err } if rec.Adapters.CLI.Enabled { profiles := make(map[string]any) for name, p := range rec.Adapters.CLI.Profiles { profiles[name] = map[string]any{ "command": p.Command, "args": stringsToAny(p.Args), "env": stringsToAny(p.Env), "persistent": p.Persistent, "terminal": p.Terminal, "response_idle_timeout_ms": p.ResponseIdleTimeoutMS, "startup_idle_timeout_ms": p.StartupIdleTimeoutMS, } } if err := addAdapter("cli", true, map[string]any{"profiles": profiles}); err != nil { return nil, err } } return payload, nil } func stringsToAny(ss []string) []interface{} { out := make([]interface{}, len(ss)) for i, s := range ss { out[i] = s } return out }