Files
lingniu-vehicle-ingest/go/vehicle-gateway/internal/gateway/tcp_server.go

452 lines
12 KiB
Go

package gateway
import (
"context"
"encoding/hex"
"errors"
"fmt"
"io"
"log/slog"
"net"
"strings"
"sync"
"time"
"lingniu-vehicle-ingest/go/vehicle-gateway/internal/envelope"
"lingniu-vehicle-ingest/go/vehicle-gateway/internal/eventbus"
"lingniu-vehicle-ingest/go/vehicle-gateway/internal/identity"
"lingniu-vehicle-ingest/go/vehicle-gateway/internal/metrics"
"lingniu-vehicle-ingest/go/vehicle-gateway/internal/realtime"
)
type FrameExtractor func([]byte) (frames [][]byte, remainder []byte, err error)
type FrameParser func(raw []byte, receivedAtMS int64, sourceEndpoint string) (envelope.FrameEnvelope, error)
type FrameResponder func(raw []byte, env envelope.FrameEnvelope) (response []byte, ok bool, err error)
type TCPProtocol struct {
Protocol envelope.Protocol
Addr string
Extract FrameExtractor
Parse FrameParser
Respond FrameResponder
}
type connectionState struct {
platformName string
}
type TCPServer struct {
protocol TCPProtocol
sink eventbus.Sink
resolver identity.Resolver
logger *slog.Logger
metrics *metrics.Registry
readBufferSize int
idleTimeout time.Duration
maxConnections int
publishUnified bool
}
type TCPServerConfig struct {
Protocol TCPProtocol
Sink eventbus.Sink
Resolver identity.Resolver
Logger *slog.Logger
Metrics *metrics.Registry
ReadBufferSize int
IdleTimeout time.Duration
MaxConnections int
PublishUnified bool
}
const frameOperationTimeout = 30 * time.Second
var gatewayFrameDurationBucketsMS = []float64{1, 5, 10, 25, 50, 100, 250, 500, 1000, 5000}
func NewTCPServer(cfg TCPServerConfig) (*TCPServer, error) {
if cfg.Protocol.Protocol == "" {
return nil, errors.New("protocol is required")
}
if strings.TrimSpace(cfg.Protocol.Addr) == "" {
return nil, errors.New("listen addr is required")
}
if cfg.Protocol.Extract == nil {
return nil, errors.New("frame extractor is required")
}
if cfg.Protocol.Parse == nil {
return nil, errors.New("frame parser is required")
}
if cfg.Sink == nil {
return nil, errors.New("sink is required")
}
if cfg.Logger == nil {
cfg.Logger = slog.Default()
}
if cfg.Resolver == nil {
cfg.Resolver = identity.NoopResolver{}
}
if cfg.ReadBufferSize <= 0 {
cfg.ReadBufferSize = 32 * 1024
}
if cfg.IdleTimeout <= 0 {
cfg.IdleTimeout = 2 * time.Minute
}
if cfg.MaxConnections <= 0 {
cfg.MaxConnections = 10_000
}
return &TCPServer{
protocol: cfg.Protocol,
sink: cfg.Sink,
resolver: cfg.Resolver,
logger: cfg.Logger,
metrics: cfg.Metrics,
readBufferSize: cfg.ReadBufferSize,
idleTimeout: cfg.IdleTimeout,
maxConnections: cfg.MaxConnections,
publishUnified: cfg.PublishUnified,
}, nil
}
func (s *TCPServer) ListenAndServe(ctx context.Context) error {
var lc net.ListenConfig
listener, err := lc.Listen(ctx, "tcp", s.protocol.Addr)
if err != nil {
return err
}
defer listener.Close()
go func() {
<-ctx.Done()
_ = listener.Close()
}()
s.logger.Info("tcp listener started", "protocol", s.protocol.Protocol, "addr", listener.Addr().String())
sem := make(chan struct{}, s.maxConnections)
var wg sync.WaitGroup
defer wg.Wait()
for {
conn, err := listener.Accept()
if err != nil {
if ctx.Err() != nil {
return nil
}
s.logger.Warn("tcp accept failed", "protocol", s.protocol.Protocol, "error", err)
continue
}
select {
case sem <- struct{}{}:
wg.Add(1)
go func() {
defer wg.Done()
defer func() { <-sem }()
s.handleConnection(ctx, conn)
}()
default:
s.recordConnectionRejection("max_connections")
s.recordConnectionClose("max_connections")
s.logger.Warn("tcp connection rejected: max connections reached", "protocol", s.protocol.Protocol, "remote", conn.RemoteAddr().String())
_ = conn.Close()
}
}
}
func (s *TCPServer) handleConnection(ctx context.Context, conn net.Conn) {
defer conn.Close()
source := conn.RemoteAddr().String()
log := s.logger.With("protocol", s.protocol.Protocol, "remote", source)
s.recordConnectionMetric(1)
defer s.recordConnectionMetric(-1)
log.Info("tcp connection opened")
defer log.Info("tcp connection closed")
readBuffer := make([]byte, s.readBufferSize)
var pending []byte
state := &connectionState{}
for {
_ = conn.SetReadDeadline(time.Now().Add(s.idleTimeout))
n, err := conn.Read(readBuffer)
if n > 0 {
pending = append(pending, readBuffer[:n]...)
frames, remainder, extractErr := s.protocol.Extract(pending)
if extractErr != nil {
log.Warn("frame extraction failed", "error", extractErr)
s.recordConnectionClose("extract_error")
return
}
pending = remainder
for _, frame := range frames {
s.handleFrame(ctx, conn, frame, source, state)
}
}
if err != nil {
if errors.Is(err, io.EOF) {
s.recordConnectionClose("eof")
return
}
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
log.Warn("tcp connection idle timeout")
s.recordConnectionClose("read_timeout")
return
}
log.Warn("tcp read failed", "error", err)
s.recordConnectionClose("read_error")
return
}
if ctx.Err() != nil {
s.recordConnectionClose("context_cancelled")
return
}
}
}
func (s *TCPServer) handleFrame(ctx context.Context, conn net.Conn, raw []byte, source string, state *connectionState) {
started := time.Now()
frameStatus := envelope.ParseBadFrame
defer func() {
s.recordFrameDuration(frameStatus, time.Since(started))
}()
frameCtx, cancelFrame := context.WithTimeout(context.WithoutCancel(ctx), frameOperationTimeout)
defer cancelFrame()
receivedAtMS := time.Now().UnixMilli()
env, err := s.protocol.Parse(raw, receivedAtMS, source)
if err != nil {
env = envelope.FrameEnvelope{
Protocol: s.protocol.Protocol,
SourceEndpoint: source,
ReceivedAtMS: receivedAtMS,
EventTimeMS: receivedAtMS,
RawHex: strings.ToUpper(hex.EncodeToString(raw)),
ParseStatus: envelope.ParseBadFrame,
ParseError: err.Error(),
}
env.EventID = env.StableEventID()
s.recordParseErrorMetric(err)
} else {
resolved, resolveErr := s.resolver.Resolve(frameCtx, env)
if resolveErr != nil {
s.logger.Warn("identity resolve failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", resolveErr)
if env.Parsed == nil {
env.Parsed = map[string]any{}
}
env.Parsed["identity"] = map[string]any{"resolved": false, "error": resolveErr.Error()}
env.ParseStatus = envelope.ParsePartial
s.recordIdentityMetric("error")
} else {
env = resolved
annotateIdentityUnresolved(&env)
s.recordIdentityMetric(identityStatus(env))
}
enrichConnectionPlatform(&env, state)
}
frameStatus = env.ParseStatus
s.recordFrameMetric(env.ParseStatus)
if env.ParseStatus != envelope.ParseBadFrame {
realtime.EnsureParsedFields(&env)
}
if err := s.sink.PublishRaw(frameCtx, env); err != nil {
s.recordPublishMetric("raw", "error")
s.logger.Error("publish raw failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", err)
return
}
s.recordPublishMetric("raw", "ok")
if env.ParseStatus == envelope.ParseBadFrame {
return
}
if fieldsEnv, ok := realtime.BuildFieldsEnvelope(env); ok {
if err := s.sink.PublishFields(frameCtx, fieldsEnv); err != nil {
s.recordPublishMetric("fields", "error")
s.logger.Error("publish fields failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", err)
return
}
s.recordPublishMetric("fields", "ok")
}
if s.publishUnified {
if err := s.sink.PublishUnified(frameCtx, env); err != nil {
s.recordPublishMetric("unified", "error")
s.logger.Error("publish unified failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", err)
return
}
s.recordPublishMetric("unified", "ok")
}
if s.protocol.Respond == nil {
return
}
response, ok, err := s.protocol.Respond(raw, env)
if err != nil {
s.logger.Warn("build protocol response failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", err)
return
}
if !ok || len(response) == 0 {
return
}
_ = conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
if _, err := conn.Write(response); err != nil {
s.logger.Warn("write protocol response failed", "protocol", s.protocol.Protocol, "event_id", env.StableEventID(), "error", err)
}
}
func enrichConnectionPlatform(env *envelope.FrameEnvelope, state *connectionState) {
if env == nil || state == nil || env.ParseStatus == envelope.ParseBadFrame {
return
}
if env.Fields == nil {
env.Fields = map[string]any{}
}
if current := strings.TrimSpace(fmt.Sprint(env.Fields["platform_account"])); current != "" && current != "<nil>" {
state.platformName = current
}
if state.platformName == "" {
return
}
if strings.TrimSpace(fmt.Sprint(env.Fields["platform_account"])) == "" || fmt.Sprint(env.Fields["platform_account"]) == "<nil>" {
env.Fields["platform_account"] = state.platformName
}
if env.Parsed == nil {
env.Parsed = map[string]any{}
}
if _, ok := env.Parsed["platform_name"]; !ok {
env.Parsed["platform_name"] = state.platformName
}
}
func (s *TCPServer) recordFrameMetric(status envelope.ParseStatus) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_frames_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"status": string(status),
})
}
func (s *TCPServer) recordFrameDuration(status envelope.ParseStatus, elapsed time.Duration) {
if s.metrics == nil {
return
}
elapsedMS := float64(elapsed.Milliseconds())
labels := metrics.Labels{
"protocol": string(s.protocol.Protocol),
"status": string(status),
}
s.metrics.SetGauge("vehicle_gateway_frame_duration_ms", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"status": string(status),
}, elapsedMS)
s.metrics.ObserveHistogram("vehicle_gateway_frame_duration_ms_histogram", labels, gatewayFrameDurationBucketsMS, elapsedMS)
}
func (s *TCPServer) recordPublishMetric(kind string, status string) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_publish_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"kind": kind,
"status": status,
})
}
func (s *TCPServer) recordParseErrorMetric(err error) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_parse_errors_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"reason": classifyError(err),
})
}
func (s *TCPServer) recordIdentityMetric(status string) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_identity_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"status": status,
})
}
func (s *TCPServer) recordConnectionMetric(delta float64) {
if s.metrics == nil {
return
}
s.metrics.AddGauge("vehicle_gateway_active_connections", metrics.Labels{
"protocol": string(s.protocol.Protocol),
}, delta)
}
func (s *TCPServer) recordConnectionRejection(reason string) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_connection_rejections_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"reason": reason,
})
}
func (s *TCPServer) recordConnectionClose(reason string) {
if s.metrics == nil {
return
}
s.metrics.IncCounter("vehicle_gateway_connection_closes_total", metrics.Labels{
"protocol": string(s.protocol.Protocol),
"reason": reason,
})
}
func identityStatus(env envelope.FrameEnvelope) string {
if strings.TrimSpace(env.VIN) != "" {
return "resolved"
}
if env.ParseStatus == envelope.ParsePartial {
return "error"
}
return "unresolved"
}
func annotateIdentityUnresolved(env *envelope.FrameEnvelope) {
if env == nil || env.ParseStatus == envelope.ParseBadFrame || strings.TrimSpace(env.VIN) != "" {
return
}
if env.Parsed == nil {
env.Parsed = map[string]any{}
}
if _, exists := env.Parsed["identity"]; exists {
return
}
env.Parsed["identity"] = map[string]any{
"resolved": false,
"reason": "no_binding",
}
}
func classifyError(err error) string {
text := strings.ToLower(strings.TrimSpace(fmt.Sprint(err)))
switch {
case text == "":
return "unknown"
case strings.Contains(text, "bcc") || strings.Contains(text, "checksum"):
return "checksum"
case strings.Contains(text, "short") || strings.Contains(text, "truncated") || strings.Contains(text, "length"):
return "length"
case strings.Contains(text, "json"):
return "json"
case strings.Contains(text, "start"):
return "start_symbol"
default:
return "parse"
}
}
func (p TCPProtocol) String() string {
return fmt.Sprintf("%s@%s", p.Protocol, p.Addr)
}