Files
go-sip/internal/rpc/server.go
T

468 lines
17 KiB
Go

package rpc
import (
"context"
cryptorand "crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"os"
"sync"
"time"
agentpb "git.ipao.vip/rogee/go-sip/gen/agent"
"git.ipao.vip/rogee/go-sip/internal/agent"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/proto"
)
// ServerOptions contains deployment-bound identity and mock policy inputs.
// Production credentials are supplied to grpc.Server through TLS credentials.
type ServerOptions struct {
Mode string
Status *agentpb.AgentStatus
Now func() time.Time
RequirePeerCertificate bool
ApprovedDispatcherID string
PeerAgentIDs map[string]string
PeerCertificateFingerprints map[string]struct{}
StatePath string
// LoadedSIP reports the revision the mock Agent actually loaded; nil fails closed.
LoadedSIP func(context.Context) (map[string]int64, error)
// The mock may have issued a call even if its outcome is unknown.
MockApprovedOriginate func(context.Context, ApprovedExecution) error
// ApprovedTaskCalls is shared with the approved call runner; nil rejects task controls.
ApprovedTaskCalls *agent.TaskCalls
}
// Server owns the current Agent session, execution and control RPC boundary.
// Dispatcher SQLite remains authoritative for quotas and task state.
type Server struct {
agentpb.UnimplementedAgentControlServiceServer
mode string
now func() time.Time
status *agentpb.AgentStatus
requirePeerCertificate bool
approvedDispatcherID string
peerAgentIDs map[string]string
peerCertificateFingerprints map[string]struct{}
loadedSIP func(context.Context) (map[string]int64, error)
mockApprovedOriginate func(context.Context, ApprovedExecution) error
approvedTaskCalls *agent.TaskCalls
sessions *SessionRegistry
approvedJournalDir string
approvedMu sync.Mutex
}
// NewServer constructs a handler suitable for registration with a gRPC server.
func NewServer(options ServerOptions) *Server {
mode := options.Mode
if mode == "" {
mode = "mock"
}
now := options.Now
if now == nil {
now = time.Now
}
statusValue := &agentpb.AgentStatus{}
if options.Status != nil {
statusValue = proto.Clone(options.Status).(*agentpb.AgentStatus)
}
if statusValue.AdmissionState == agentpb.AdmissionState_ADMISSION_STATE_UNSPECIFIED {
statusValue.AdmissionState = agentpb.AdmissionState_ADMISSION_STATE_CLOSED
}
server := &Server{
mode: mode,
now: now,
status: statusValue,
requirePeerCertificate: options.RequirePeerCertificate,
approvedDispatcherID: options.ApprovedDispatcherID,
peerAgentIDs: cloneStringMap(options.PeerAgentIDs),
peerCertificateFingerprints: cloneSet(options.PeerCertificateFingerprints),
loadedSIP: options.LoadedSIP,
mockApprovedOriginate: options.MockApprovedOriginate,
approvedTaskCalls: options.ApprovedTaskCalls,
sessions: NewSessionRegistry(options.StatePath),
}
if options.StatePath != "" {
server.approvedJournalDir = options.StatePath + ".approved"
}
return server
}
// SessionRegistry keeps the newest Dispatcher-approved binding for each Agent.
// A newer generation fences all older requests; it does not release unknown
// work from an older boot.
type SessionRegistry struct {
mu sync.Mutex
sessions map[string]sessionRecord
generations map[string]uint64
statePath string
loadErr error
}
type sessionRecord struct {
binding *agentpb.AgentBinding
activationOperationID string
digest string
session *agentpb.Session
}
func NewSessionRegistry(statePath string) *SessionRegistry {
registry := &SessionRegistry{sessions: make(map[string]sessionRecord), generations: make(map[string]uint64), statePath: statePath}
if statePath != "" {
registry.loadErr = registry.load()
}
return registry
}
type sessionJournal struct {
Generations map[string]uint64 `json:"generations"`
}
func (r *SessionRegistry) load() error {
data, err := os.ReadFile(r.statePath)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return err
}
var journal sessionJournal
if err := json.Unmarshal(data, &journal); err != nil {
return err
}
for agentID, generation := range journal.Generations {
if agentID != "" && generation > 0 {
r.generations[agentID] = generation
}
}
return nil
}
func (r *SessionRegistry) persistLocked() error {
if r.statePath == "" {
return nil
}
return writeRPCJournal(r.statePath, sessionJournal{Generations: r.generations})
}
func (r *SessionRegistry) Activate(binding *agentpb.AgentBinding, activationOperationID, digest string, now time.Time) (*agentpb.Session, bool, error) {
if binding == nil || binding.AgentId == "" || binding.CellId == "" || binding.DispatcherEpoch == "" || activationOperationID == "" {
return nil, false, status.Error(codes.InvalidArgument, "agent, cell, epoch and activation operation are required")
}
r.mu.Lock()
defer r.mu.Unlock()
if r.loadErr != nil {
return nil, false, status.Errorf(codes.Internal, "load session journal: %v", r.loadErr)
}
if existing, ok := r.sessions[binding.AgentId]; ok {
if existing.activationOperationID == activationOperationID && existing.digest == digest {
return cloneSession(existing.session), true, nil
}
if binding.SessionGeneration == 0 {
if existing.binding.SessionGeneration == ^uint64(0) {
return nil, false, status.Error(codes.Aborted, "session generation is exhausted")
}
binding = proto.Clone(binding).(*agentpb.AgentBinding)
binding.SessionGeneration = existing.binding.SessionGeneration + 1
}
if binding.SessionGeneration <= existing.binding.SessionGeneration {
return nil, false, status.Error(codes.Aborted, "session generation is fenced")
}
}
if binding.SessionGeneration == 0 {
binding = proto.Clone(binding).(*agentpb.AgentBinding)
binding.SessionGeneration = 1
if previous, ok := r.generations[binding.AgentId]; ok {
if previous == ^uint64(0) {
return nil, false, status.Error(codes.Aborted, "session generation is exhausted")
}
binding.SessionGeneration = previous + 1
}
}
if previous, ok := r.generations[binding.AgentId]; ok && binding.SessionGeneration <= previous {
return nil, false, status.Error(codes.Aborted, "persisted session generation is fenced")
}
credential := make([]byte, 32)
if _, err := cryptorand.Read(credential); err != nil {
return nil, false, status.Errorf(codes.Internal, "create session credential: %v", err)
}
session := &agentpb.Session{
DispatcherEpoch: binding.DispatcherEpoch,
SessionGeneration: binding.SessionGeneration,
ExpiresAtUnixMs: now.Add(10 * time.Minute).UnixMilli(),
SessionCredential: credential,
}
previous, hadPrevious := r.sessions[binding.AgentId]
r.sessions[binding.AgentId] = sessionRecord{
binding: proto.Clone(binding).(*agentpb.AgentBinding),
activationOperationID: activationOperationID,
digest: digest,
session: cloneSession(session),
}
r.generations[binding.AgentId] = binding.SessionGeneration
if err := r.persistLocked(); err != nil {
if hadPrevious {
r.sessions[binding.AgentId] = previous
} else {
delete(r.sessions, binding.AgentId)
}
return nil, false, status.Errorf(codes.Internal, "persist session journal: %v", err)
}
return session, false, nil
}
func (r *SessionRegistry) Authorize(meta *agentpb.RequestMeta, now time.Time) error {
if meta == nil || meta.AgentId == "" || meta.CellId == "" || meta.BootId == "" || meta.DispatcherEpoch == "" || meta.SessionGeneration == 0 {
return status.Error(codes.InvalidArgument, "complete session metadata is required")
}
r.mu.Lock()
defer r.mu.Unlock()
existing, ok := r.sessions[meta.AgentId]
if !ok {
return status.Error(codes.Unauthenticated, "agent session is not active")
}
if existing.binding.CellId != meta.CellId || existing.binding.ExpectedBootId != meta.BootId || existing.binding.DispatcherEpoch != meta.DispatcherEpoch || existing.binding.SessionGeneration != meta.SessionGeneration {
return status.Error(codes.Aborted, "agent session is fenced")
}
if existing.session.ExpiresAtUnixMs <= now.UnixMilli() {
return status.Error(codes.Unauthenticated, "agent session expired")
}
return nil
}
// ApprovedDispatcher returns the Dispatcher ID attached to the authenticated
// active session, not the ID asserted by an individual execution request.
func (r *SessionRegistry) ApprovedDispatcher(meta *agentpb.RequestMeta, now time.Time) (string, error) {
if err := r.Authorize(meta, now); err != nil {
return "", err
}
r.mu.Lock()
defer r.mu.Unlock()
session, ok := r.sessions[meta.AgentId]
if !ok || session.binding.SessionGeneration != meta.SessionGeneration || session.session.ExpiresAtUnixMs <= now.UnixMilli() {
return "", status.Error(codes.Aborted, "agent session changed")
}
if session.binding.DispatcherId == "" {
return "", status.Error(codes.FailedPrecondition, "active Dispatcher identity is missing")
}
return session.binding.DispatcherId, nil
}
// CurrentMeta returns only the active, unexpired session identity. It does
// not expose the session credential; callers use it for Agent→Dispatcher
// reports and never invent a new boot identity after restart.
func (r *SessionRegistry) CurrentMeta(agentID string, now time.Time) (*agentpb.RequestMeta, error) {
if agentID == "" {
return nil, status.Error(codes.InvalidArgument, "Agent ID is required")
}
r.mu.Lock()
defer r.mu.Unlock()
if r.loadErr != nil {
return nil, status.Errorf(codes.Internal, "load session journal: %v", r.loadErr)
}
existing, ok := r.sessions[agentID]
if !ok || existing.session.ExpiresAtUnixMs <= now.UnixMilli() {
return nil, status.Error(codes.Unauthenticated, "Agent session is not active")
}
return &agentpb.RequestMeta{
ProtocolVersion: "agent.v1", AgentId: agentID, CellId: existing.binding.CellId,
BootId: existing.binding.ExpectedBootId, DispatcherEpoch: existing.binding.DispatcherEpoch,
SessionGeneration: existing.binding.SessionGeneration,
}, nil
}
func (s *Server) ActiveSessionMeta() (*agentpb.RequestMeta, error) {
if s.status == nil {
return nil, status.Error(codes.FailedPrecondition, "Agent status is unavailable")
}
return s.sessions.CurrentMeta(s.status.AgentId, s.now())
}
func (s *Server) GetAgentStatus(ctx context.Context, req *agentpb.GetAgentStatusRequest) (*agentpb.GetAgentStatusResponse, error) {
if req == nil || req.Meta == nil || req.Meta.AgentId == "" || req.Meta.CellId == "" {
return nil, status.Error(codes.InvalidArgument, "status metadata with agent and cell is required")
}
if err := s.checkConfiguredIdentity(req.Meta.AgentId, req.Meta.CellId); err != nil {
return nil, err
}
if req.Target != nil {
if req.Target.AgentId != "" && req.Target.AgentId != req.Meta.AgentId {
return nil, status.Error(codes.PermissionDenied, "target agent does not match authenticated agent")
}
if req.Target.CellId != "" && req.Target.CellId != req.Meta.CellId {
return nil, status.Error(codes.PermissionDenied, "target Cell does not match authenticated Cell")
}
}
if err := s.checkPeer(ctx, req.Meta.AgentId); err != nil {
return nil, err
}
preActivation := req.Meta.BootId == "" && req.Meta.DispatcherEpoch == "" && req.Meta.SessionGeneration == 0
if preActivation {
if req.Target != nil && (req.Target.ExpectedBootId != "" || req.Target.DispatcherEpoch != "" || req.Target.SessionGeneration != 0) {
return nil, status.Error(codes.InvalidArgument, "pre-activation status cannot include session binding")
}
} else if err := s.sessions.Authorize(req.Meta, s.now()); err != nil {
return nil, err
}
result := proto.Clone(s.status).(*agentpb.AgentStatus)
result.SessionActive = !preActivation
result.MtlsAuthenticated = s.peerIsAuthenticated(ctx)
return &agentpb.GetAgentStatusResponse{Meta: s.responseMeta(req.Meta), Status: result}, nil
}
func (s *Server) ActivateAgent(ctx context.Context, req *agentpb.ActivateAgentRequest) (*agentpb.ActivateAgentResponse, error) {
if req == nil || req.Meta == nil || req.Binding == nil {
return nil, status.Error(codes.InvalidArgument, "activation metadata and binding are required")
}
if req.Meta.AgentId != req.Binding.AgentId || req.Meta.CellId != req.Binding.CellId || req.Meta.BootId == "" || req.Meta.DispatcherEpoch == "" {
return nil, status.Error(codes.InvalidArgument, "activation identity is inconsistent")
}
if err := s.checkConfiguredIdentity(req.Binding.AgentId, req.Binding.CellId); err != nil {
return nil, err
}
if err := s.checkPeer(ctx, req.Binding.AgentId); err != nil {
return nil, err
}
if s.approvedDispatcherID != "" && req.Binding.DispatcherId != s.approvedDispatcherID {
return nil, status.Error(codes.PermissionDenied, "activated Dispatcher identity is not authorized")
}
binding := proto.Clone(req.Binding).(*agentpb.AgentBinding)
if binding.ExpectedBootId == "" {
binding.ExpectedBootId = req.Meta.BootId
}
if binding.ExpectedBootId != req.Meta.BootId {
return nil, status.Error(codes.Aborted, "activation boot identity is fenced")
}
if req.ActivationOperationId == "" {
return nil, status.Error(codes.InvalidArgument, "activation operation is required")
}
digest := messageDigest(req)
session, replay, err := s.sessions.Activate(binding, req.ActivationOperationId, digest, s.now())
if err != nil {
return nil, err
}
state := agentpb.ActivationState_ACTIVATION_STATE_ACTIVE
if replay {
state = agentpb.ActivationState_ACTIVATION_STATE_ACTIVE
}
return &agentpb.ActivateAgentResponse{Meta: s.responseMeta(req.Meta), State: state, Session: session}, nil
}
func (s *Server) authorize(ctx context.Context, meta *agentpb.RequestMeta) error {
if meta == nil {
return status.Error(codes.InvalidArgument, "request metadata is required")
}
if err := s.checkPeer(ctx, meta.AgentId); err != nil {
return err
}
return s.sessions.Authorize(meta, s.now())
}
func (s *Server) checkConfiguredIdentity(agentID, cellID string) error {
if s.status.AgentId != "" && s.status.AgentId != agentID {
return status.Error(codes.PermissionDenied, "request Agent identity is not bound to this endpoint")
}
if s.status.CellId != "" && s.status.CellId != cellID {
return status.Error(codes.PermissionDenied, "request Cell identity is not bound to this endpoint")
}
return nil
}
func (s *Server) checkPeer(ctx context.Context, agentID string) error {
if !s.requirePeerCertificate && len(s.peerAgentIDs) == 0 && len(s.peerCertificateFingerprints) == 0 {
return nil
}
p, ok := peer.FromContext(ctx)
if !ok {
return status.Error(codes.Unauthenticated, "mTLS peer is missing")
}
tlsInfo, ok := p.AuthInfo.(credentials.TLSInfo)
if !ok || len(tlsInfo.State.VerifiedChains) == 0 || len(tlsInfo.State.VerifiedChains[0]) == 0 {
return status.Error(codes.Unauthenticated, "verified mTLS peer is required")
}
fingerprint := CertificateFingerprint(tlsInfo.State.VerifiedChains[0][0])
if len(s.peerAgentIDs) == 0 && len(s.peerCertificateFingerprints) == 0 {
// The shared Agent certificate authenticates the certificate group. The
// Dispatcher-approved session binding still authorizes the individual
// agent/cell/boot tuple; no self-reported identity is trusted here.
return nil
}
if len(s.peerCertificateFingerprints) > 0 {
if _, allowed := s.peerCertificateFingerprints[fingerprint]; !allowed {
return status.Error(codes.PermissionDenied, "mTLS certificate is not in the endpoint allowlist")
}
}
if len(s.peerAgentIDs) > 0 {
if expected := s.peerAgentIDs[fingerprint]; expected == "" || expected != agentID {
return status.Error(codes.PermissionDenied, "mTLS certificate is not bound to this agent")
}
}
return nil
}
func (s *Server) peerIsAuthenticated(ctx context.Context) bool {
p, ok := peer.FromContext(ctx)
if !ok {
return false
}
_, ok = p.AuthInfo.(credentials.TLSInfo)
return ok
}
func (s *Server) responseMeta(meta *agentpb.RequestMeta) *agentpb.ResponseMeta {
if meta == nil {
return nil
}
return &agentpb.ResponseMeta{ProtocolVersion: meta.ProtocolVersion, RequestId: meta.RequestId, TraceId: meta.TraceId, OperationId: meta.OperationId, ObservedAtUnixMs: s.now().UnixMilli(), DispatcherEpoch: meta.DispatcherEpoch, AgentId: meta.AgentId, CellId: meta.CellId, BootId: meta.BootId, SessionGeneration: meta.SessionGeneration}
}
func messageDigest(message proto.Message) string {
encoded, err := proto.Marshal(message)
if err != nil {
return "marshal-error"
}
digest := sha256.Sum256(encoded)
return hex.EncodeToString(digest[:])
}
func cloneSession(value *agentpb.Session) *agentpb.Session {
if value == nil {
return nil
}
return proto.Clone(value).(*agentpb.Session)
}
func cloneStringMap(values map[string]string) map[string]string {
if values == nil {
return nil
}
result := make(map[string]string, len(values))
for key, value := range values {
result[key] = value
}
return result
}
func cloneSet(values map[string]struct{}) map[string]struct{} {
if values == nil {
return nil
}
result := make(map[string]struct{}, len(values))
for key := range values {
result[key] = struct{}{}
}
return result
}
var _ agentpb.AgentControlServiceServer = (*Server)(nil)
var _ = grpc.SupportPackageIsVersion9