iop/apps/edge/internal/transport/server.go

216 lines
5.7 KiB
Go

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
}