amaniquery-agent / pkg /observability /observability.go
Deployment
Automated deployment update
4b1daed
Raw
History Blame Contribute Delete
8.42 kB
// Package observability provides OpenTelemetry tracing, Prometheus metrics, and structured logging
package observability
import (
"context"
"fmt"
"net/http"
"time"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
"github.com/prometheus/client_golang/prometheus/promhttp"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc"
"go.opentelemetry.io/otel/propagation"
"go.opentelemetry.io/otel/sdk/resource"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.21.0"
"go.opentelemetry.io/otel/trace"
"go.uber.org/zap"
"google.golang.org/grpc"
"google.golang.org/grpc/status"
)
// Config for observability
type Config struct {
ServiceName string
ServiceVersion string
TracingEnabled bool
TracingEndpoint string
MetricsEnabled bool
MetricsPort int
}
// Metrics for the RAG service
var (
queryCounter = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "amaniquery_queries_total",
Help: "Total number of queries processed",
},
[]string{"status", "cache_hit", "strategy"},
)
queryDuration = promauto.NewHistogramVec(
prometheus.HistogramOpts{
Name: "amaniquery_query_duration_seconds",
Help: "Query processing duration in seconds",
Buckets: []float64{0.1, 0.25, 0.5, 1, 2.5, 5, 10},
},
[]string{"strategy"},
)
retrievalDuration = promauto.NewHistogramVec(
prometheus.HistogramOpts{
Name: "amaniquery_retrieval_duration_seconds",
Help: "Document retrieval duration in seconds",
Buckets: []float64{0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1},
},
[]string{"type"},
)
generationDuration = promauto.NewHistogram(
prometheus.HistogramOpts{
Name: "amaniquery_generation_duration_seconds",
Help: "LLM generation duration in seconds",
Buckets: []float64{0.5, 1, 2.5, 5, 10, 30},
},
)
tokensUsed = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "amaniquery_tokens_total",
Help: "Total tokens used",
},
[]string{"type"},
)
cacheHits = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "amaniquery_cache_hits_total",
Help: "Cache hit count",
},
[]string{"tier"},
)
cacheMisses = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "amaniquery_cache_misses_total",
Help: "Cache miss count",
},
[]string{"tier"},
)
documentsIndexed = promauto.NewCounter(
prometheus.CounterOpts{
Name: "amaniquery_documents_indexed_total",
Help: "Total documents indexed",
},
)
activeConnections = promauto.NewGauge(
prometheus.GaugeOpts{
Name: "amaniquery_active_connections",
Help: "Number of active connections",
},
)
errorCounter = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "amaniquery_errors_total",
Help: "Total errors by type",
},
[]string{"type", "component"},
)
)
// InitProvider initializes the OpenTelemetry tracer provider
func InitProvider(cfg Config) (func(context.Context) error, error) {
if !cfg.TracingEnabled {
return func(ctx context.Context) error { return nil }, nil
}
ctx := context.Background()
// Create resource
res, err := resource.Merge(
resource.Default(),
resource.NewWithAttributes(
semconv.SchemaURL,
semconv.ServiceName(cfg.ServiceName),
semconv.ServiceVersion(cfg.ServiceVersion),
),
)
if err != nil {
return nil, err
}
// Create OTLP exporter
exporter, err := otlptracegrpc.New(ctx,
otlptracegrpc.WithEndpoint(cfg.TracingEndpoint),
otlptracegrpc.WithInsecure(),
)
if err != nil {
return nil, err
}
// Create tracer provider
tp := sdktrace.NewTracerProvider(
sdktrace.WithBatcher(exporter),
sdktrace.WithResource(res),
sdktrace.WithSampler(sdktrace.AlwaysSample()),
)
// Set global tracer provider
otel.SetTracerProvider(tp)
otel.SetTextMapPropagator(propagation.NewCompositeTextMapPropagator(
propagation.TraceContext{},
propagation.Baggage{},
))
return tp.Shutdown, nil
}
// StartMetricsServer starts the Prometheus metrics server
func StartMetricsServer(port int) *http.Server {
mux := http.NewServeMux()
mux.Handle("/metrics", promhttp.Handler())
mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.Write([]byte("OK"))
})
server := &http.Server{
Addr: fmt.Sprintf(":%d", port),
Handler: mux,
}
go func() {
if err := server.ListenAndServe(); err != http.ErrServerClosed {
// Log error
}
}()
return server
}
// RecordQuery records query metrics
func RecordQuery(status string, cacheHit bool, strategy string, duration time.Duration) {
cacheHitStr := "false"
if cacheHit {
cacheHitStr = "true"
}
queryCounter.WithLabelValues(status, cacheHitStr, strategy).Inc()
queryDuration.WithLabelValues(strategy).Observe(duration.Seconds())
}
// RecordRetrieval records retrieval metrics
func RecordRetrieval(retrievalType string, duration time.Duration) {
retrievalDuration.WithLabelValues(retrievalType).Observe(duration.Seconds())
}
// RecordGeneration records LLM generation metrics
func RecordGeneration(duration time.Duration) {
generationDuration.Observe(duration.Seconds())
}
// RecordTokens records token usage
func RecordTokens(tokenType string, count int) {
tokensUsed.WithLabelValues(tokenType).Add(float64(count))
}
// RecordCacheHit records a cache hit
func RecordCacheHit(tier string) {
cacheHits.WithLabelValues(tier).Inc()
}
// RecordCacheMiss records a cache miss
func RecordCacheMiss(tier string) {
cacheMisses.WithLabelValues(tier).Inc()
}
// RecordDocumentIndexed records a document being indexed
func RecordDocumentIndexed() {
documentsIndexed.Inc()
}
// RecordError records an error
func RecordError(errorType, component string) {
errorCounter.WithLabelValues(errorType, component).Inc()
}
// IncrementConnections increments active connections
func IncrementConnections() {
activeConnections.Inc()
}
// DecrementConnections decrements active connections
func DecrementConnections() {
activeConnections.Dec()
}
// UnaryServerInterceptor returns a gRPC unary server interceptor for tracing
func UnaryServerInterceptor() grpc.UnaryServerInterceptor {
return func(
ctx context.Context,
req interface{},
info *grpc.UnaryServerInfo,
handler grpc.UnaryHandler,
) (interface{}, error) {
tracer := otel.Tracer("grpc-server")
ctx, span := tracer.Start(ctx, info.FullMethod,
trace.WithSpanKind(trace.SpanKindServer),
)
defer span.End()
start := time.Now()
resp, err := handler(ctx, req)
duration := time.Since(start)
if err != nil {
span.SetAttributes(attribute.String("error", err.Error()))
st, _ := status.FromError(err)
span.SetAttributes(attribute.String("grpc.status_code", st.Code().String()))
}
span.SetAttributes(
attribute.String("grpc.method", info.FullMethod),
attribute.Int64("grpc.duration_ms", duration.Milliseconds()),
)
return resp, err
}
}
// StreamServerInterceptor returns a gRPC stream server interceptor for tracing
func StreamServerInterceptor() grpc.StreamServerInterceptor {
return func(
srv interface{},
ss grpc.ServerStream,
info *grpc.StreamServerInfo,
handler grpc.StreamHandler,
) error {
tracer := otel.Tracer("grpc-server")
ctx, span := tracer.Start(ss.Context(), info.FullMethod,
trace.WithSpanKind(trace.SpanKindServer),
)
defer span.End()
wrappedStream := &tracedServerStream{
ServerStream: ss,
ctx: ctx,
}
err := handler(srv, wrappedStream)
if err != nil {
span.SetAttributes(attribute.String("error", err.Error()))
}
return err
}
}
type tracedServerStream struct {
grpc.ServerStream
ctx context.Context
}
func (s *tracedServerStream) Context() context.Context {
return s.ctx
}
// Logger creates a structured logger
func NewLogger(level, format string) (*zap.Logger, error) {
var config zap.Config
if format == "json" {
config = zap.NewProductionConfig()
} else {
config = zap.NewDevelopmentConfig()
}
switch level {
case "debug":
config.Level = zap.NewAtomicLevelAt(zap.DebugLevel)
case "info":
config.Level = zap.NewAtomicLevelAt(zap.InfoLevel)
case "warn":
config.Level = zap.NewAtomicLevelAt(zap.WarnLevel)
case "error":
config.Level = zap.NewAtomicLevelAt(zap.ErrorLevel)
default:
config.Level = zap.NewAtomicLevelAt(zap.InfoLevel)
}
return config.Build()
}