Deployment
Automated deployment update
4b1daed
Raw
History Blame Contribute Delete
7.81 kB
// Package main provides the entry point for the Generator service.
// This service handles LLM interactions for text generation and embeddings.
package main
import (
"context"
"fmt"
"net"
"os"
"os/signal"
"syscall"
"github.com/AmaniQuery/amaniquery/internal/generator"
"github.com/AmaniQuery/amaniquery/internal/generator/llm"
"github.com/AmaniQuery/amaniquery/pkg/config"
"github.com/AmaniQuery/amaniquery/pkg/observability"
"go.uber.org/zap"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/health"
"google.golang.org/grpc/health/grpc_health_v1"
"google.golang.org/grpc/reflection"
"google.golang.org/grpc/status"
)
func main() {
// Load configuration
cfg, err := config.Load()
if err != nil {
fmt.Fprintf(os.Stderr, "failed to load configuration: %v\n", err)
os.Exit(1)
}
// Initialize logger
logger, err := observability.NewLogger(cfg.Observability.LogLevel, cfg.Observability.LogFormat)
if err != nil {
fmt.Fprintf(os.Stderr, "failed to initialize logger: %v\n", err)
os.Exit(1)
}
defer logger.Sync()
logger.Info("starting AmaniQuery generator server",
zap.String("version", cfg.Version),
zap.Strings("llm_providers", cfg.LLM.GetConfiguredProviders()),
)
// Initialize observability
tracingShutdown, err := observability.InitProvider(observability.Config{
ServiceName: "amaniquery-generator",
ServiceVersion: cfg.Version,
TracingEnabled: cfg.Observability.TracingEnabled,
TracingEndpoint: cfg.Observability.TracingEndpoint,
})
if err != nil {
logger.Warn("failed to initialize tracing", zap.Error(err))
} else {
defer tracingShutdown(context.Background())
}
// Build dependencies
deps, err := buildGeneratorDependencies(cfg, logger)
if err != nil {
logger.Fatal("failed to build dependencies", zap.Error(err))
}
// Create gRPC server
grpcServer := grpc.NewServer(
grpc.ChainUnaryInterceptor(
observability.UnaryServerInterceptor(),
),
)
// Register generator service
generatorServer := NewGeneratorServer(deps, logger)
RegisterGeneratorServiceServer(grpcServer, generatorServer)
// Register health service
healthServer := health.NewServer()
grpc_health_v1.RegisterHealthServer(grpcServer, healthServer)
healthServer.SetServingStatus("", grpc_health_v1.HealthCheckResponse_SERVING)
// Enable reflection
reflection.Register(grpcServer)
// Start server
port := 9092 // Different port from other servers
addr := fmt.Sprintf(":%d", port)
listener, err := net.Listen("tcp", addr)
if err != nil {
logger.Fatal("failed to listen", zap.String("addr", addr), zap.Error(err))
}
logger.Info("gRPC generator server starting", zap.String("addr", addr))
// Graceful shutdown
errChan := make(chan error, 1)
go func() {
errChan <- grpcServer.Serve(listener)
}()
quit := make(chan os.Signal, 1)
signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
select {
case err := <-errChan:
logger.Fatal("server error", zap.Error(err))
case sig := <-quit:
logger.Info("shutting down", zap.String("signal", sig.String()))
}
grpcServer.GracefulStop()
logger.Info("generator server stopped")
}
// GeneratorDependencies holds generator service dependencies
type GeneratorDependencies struct {
LLMClient *llm.FallbackClient
EmbeddingClient *generator.OpenAIEmbeddingClient
}
func buildGeneratorDependencies(cfg *config.Config, logger *zap.Logger) (*GeneratorDependencies, error) {
// Initialize LLM client with fallback
llmClient := llm.NewFallbackClient(llm.Config{
GeminiAPIKey: cfg.LLM.GeminiAPIKey,
MoonshotAPIKey: cfg.LLM.MoonshotAPIKey,
OllamaBaseURL: cfg.LLM.OllamaBaseURL,
OpenAIAPIKey: cfg.LLM.OpenAIAPIKey,
AnthropicAPIKey: cfg.LLM.AnthropicAPIKey,
DefaultModel: cfg.LLM.DefaultModel,
MaxTokens: cfg.LLM.MaxTokens,
Temperature: cfg.LLM.Temperature,
Timeout: cfg.LLM.Timeout,
MaxRetries: cfg.LLM.MaxRetries,
EnableFallback: cfg.LLM.EnableFallback,
Logger: logger,
})
// Initialize embedding client
embeddingClient := generator.NewOpenAIEmbeddingClient(generator.EmbeddingConfig{
Provider: cfg.Embedding.Provider,
APIKey: cfg.Embedding.APIKey,
Model: cfg.Embedding.Model,
Dimension: cfg.Embedding.Dimension,
BatchSize: cfg.Embedding.BatchSize,
})
return &GeneratorDependencies{
LLMClient: llmClient,
EmbeddingClient: embeddingClient,
}, nil
}
// GeneratorServer implements the GeneratorService gRPC server
type GeneratorServer struct {
UnimplementedGeneratorServiceServer
deps *GeneratorDependencies
logger *zap.Logger
}
// NewGeneratorServer creates a new generator server
func NewGeneratorServer(deps *GeneratorDependencies, logger *zap.Logger) *GeneratorServer {
return &GeneratorServer{
deps: deps,
logger: logger,
}
}
// Generate performs text generation
func (s *GeneratorServer) Generate(ctx context.Context, req *GenerateRequest) (*GenerateResponse, error) {
if len(req.Messages) == 0 {
return nil, status.Error(codes.InvalidArgument, "messages are required")
}
// Convert messages
messages := make([]llm.Message, len(req.Messages))
for i, m := range req.Messages {
messages[i] = llm.Message{
Role: m.Role,
Content: m.Content,
}
}
// Generate response
resp, err := s.deps.LLMClient.Generate(ctx, messages, llm.Options{
Model: req.Model,
Temperature: req.Temperature,
MaxTokens: int(req.MaxTokens),
})
if err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
return &GenerateResponse{
Content: resp.Content,
FinishReason: resp.FinishReason,
Model: resp.Model,
Provider: string(resp.Provider),
Usage: &TokenUsage{
PromptTokens: int32(resp.Usage.PromptTokens),
CompletionTokens: int32(resp.Usage.CompletionTokens),
TotalTokens: int32(resp.Usage.TotalTokens),
},
}, nil
}
// GenerateEmbedding generates embeddings for text
func (s *GeneratorServer) GenerateEmbedding(ctx context.Context, req *EmbeddingRequest) (*EmbeddingResponse, error) {
if req.Text == "" {
return nil, status.Error(codes.InvalidArgument, "text is required")
}
embedding, err := s.deps.EmbeddingClient.Generate(ctx, req.Text)
if err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
return &EmbeddingResponse{
Embedding: embedding,
Dimension: int32(len(embedding)),
}, nil
}
// BatchGenerateEmbeddings generates embeddings for multiple texts
func (s *GeneratorServer) BatchGenerateEmbeddings(ctx context.Context, req *BatchEmbeddingRequest) (*BatchEmbeddingResponse, error) {
if len(req.Texts) == 0 {
return nil, status.Error(codes.InvalidArgument, "texts are required")
}
embeddings, err := s.deps.EmbeddingClient.GenerateBatch(ctx, req.Texts)
if err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
return &BatchEmbeddingResponse{
Embeddings: embeddings,
}, nil
}
// Stub types - replace with generated protobuf code
type Message struct {
Role string
Content string
}
type GenerateRequest struct {
Messages []*Message
Model string
Temperature float32
MaxTokens int32
}
type GenerateResponse struct {
Content string
FinishReason string
Model string
Provider string
Usage *TokenUsage
}
type TokenUsage struct {
PromptTokens int32
CompletionTokens int32
TotalTokens int32
}
type EmbeddingRequest struct {
Text string
Model string
}
type EmbeddingResponse struct {
Embedding []float32
Dimension int32
}
type BatchEmbeddingRequest struct {
Texts []string
Model string
}
type BatchEmbeddingResponse struct {
Embeddings [][]float32
}
type UnimplementedGeneratorServiceServer struct{}
func RegisterGeneratorServiceServer(s *grpc.Server, srv *GeneratorServer) {
// Registration happens when protobuf is generated
}