| package server |
|
|
| import ( |
| "context" |
| "errors" |
| "fmt" |
| "log/slog" |
| "net/http" |
|
|
| "github.com/getkin/kin-openapi/openapi3filter" |
| "github.com/go-chi/chi/v5" |
| "github.com/go-chi/chi/v5/middleware" |
| "github.com/go-chi/cors" |
| "github.com/go-chi/render" |
| oapimiddleware "github.com/oapi-codegen/nethttp-middleware" |
| "github.com/samber/lo" |
|
|
| "github.com/openmeterio/openmeter/api" |
| v3server "github.com/openmeterio/openmeter/api/v3/server" |
| appconfig "github.com/openmeterio/openmeter/app/config" |
| "github.com/openmeterio/openmeter/openmeter/portal/authenticator" |
| "github.com/openmeterio/openmeter/openmeter/server/router" |
| "github.com/openmeterio/openmeter/pkg/contextx" |
| "github.com/openmeterio/openmeter/pkg/models" |
| "github.com/openmeterio/openmeter/pkg/server" |
| ) |
|
|
| type Server struct { |
| chi.Router |
| } |
|
|
| type ServerLogger struct{} |
|
|
| type MiddlewareManager interface { |
| Use(middlewares ...func(http.Handler) http.Handler) |
| } |
|
|
| type MiddlewareHook func(m MiddlewareManager) |
|
|
| type RouteManager interface { |
| Mount(pattern string, h http.Handler) |
| Handle(pattern string, h http.Handler) |
| HandleFunc(pattern string, h http.HandlerFunc) |
| Method(method, pattern string, h http.Handler) |
| MethodFunc(method, pattern string, h http.HandlerFunc) |
| Connect(pattern string, h http.HandlerFunc) |
| Delete(pattern string, h http.HandlerFunc) |
| Get(pattern string, h http.HandlerFunc) |
| Head(pattern string, h http.HandlerFunc) |
| Options(pattern string, h http.HandlerFunc) |
| Patch(pattern string, h http.HandlerFunc) |
| Post(pattern string, h http.HandlerFunc) |
| Put(pattern string, h http.HandlerFunc) |
| Trace(pattern string, h http.HandlerFunc) |
| } |
|
|
| type RouteHook func(r RouteManager) |
|
|
| type RouterHooks struct { |
| Middlewares []MiddlewareHook |
| Routes []RouteHook |
| } |
|
|
| type PostAuthMiddlewares []server.MiddlewareFunc |
|
|
| var _ models.Validator = (*Config)(nil) |
|
|
| type Config struct { |
| RouterConfig router.Config |
| RouterHooks RouterHooks |
| PostAuthMiddlewares PostAuthMiddlewares |
| ResponseValidation appconfig.ResponseValidationConfig |
| ClientIPMiddleware server.MiddlewareFunc |
| } |
|
|
| func (c Config) Validate() error { |
| var errs []error |
|
|
| if err := c.RouterConfig.Validate(); err != nil { |
| errs = append(errs, fmt.Errorf("invalid router config: %w", err)) |
| } |
|
|
| if c.ClientIPMiddleware == nil { |
| errs = append(errs, errors.New("client IP middleware is required")) |
| } |
|
|
| return errors.Join(errs...) |
| } |
|
|
| func NewServer(config *Config) (*Server, error) { |
| if err := config.Validate(); err != nil { |
| return nil, fmt.Errorf("invalid server config: %w", err) |
| } |
|
|
| |
| swagger, err := api.GetSwagger() |
| if err != nil { |
| return nil, fmt.Errorf("failed to get swagger: %w", err) |
| } |
|
|
| |
| |
| swagger.Servers = nil |
|
|
| impl, err := router.NewRouter(config.RouterConfig) |
| if err != nil { |
| return nil, fmt.Errorf("failed to create API: %w", err) |
| } |
|
|
| r := chi.NewRouter() |
| r.Use(server.NewPoweredByMiddleware()) |
|
|
| |
| |
| |
| |
| hookMiddlewares := collectMiddlewareHooks(config.RouterHooks.Middlewares) |
|
|
| |
| |
| v3Middlewares := append([]server.MiddlewareFunc{}, hookMiddlewares...) |
| v3Middlewares = append(v3Middlewares, []server.MiddlewareFunc{ |
| config.ClientIPMiddleware, |
| middleware.RequestID, |
| func(h http.Handler) http.Handler { |
| return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| ctx := r.Context() |
| ctx = contextx.WithAttrs(ctx, server.GetRequestAttributes(r)) |
|
|
| h.ServeHTTP(w, r.WithContext(ctx)) |
| }) |
| }, |
| server.NewRequestLoggerMiddleware(slog.Default().Handler()), |
| middleware.Recoverer, |
| }...) |
|
|
| v3API, err := v3server.NewServer(&v3server.Config{ |
| BaseURL: "/api/v3", |
| NamespaceDecoder: config.RouterConfig.NamespaceDecoder, |
| ErrorHandler: config.RouterConfig.ErrorHandler, |
| Credits: config.RouterConfig.Credits, |
| UnitConfig: config.RouterConfig.UnitConfig, |
| AddonService: config.RouterConfig.Addon, |
| AppService: config.RouterConfig.App, |
| BillingService: config.RouterConfig.Billing, |
| CustomerService: config.RouterConfig.Customer, |
| CreditGrantService: config.RouterConfig.CreditGrantService, |
| Ledger: config.RouterConfig.Ledger, |
| AccountResolver: config.RouterConfig.AccountResolver, |
| CustomerBalanceFacade: config.RouterConfig.CustomerBalanceFacade, |
| CurrencyService: config.RouterConfig.CurrencyService, |
| EntitlementService: config.RouterConfig.EntitlementConnector, |
| GovernanceService: config.RouterConfig.GovernanceService, |
| IngestService: config.RouterConfig.IngestService, |
| MeterEventService: config.RouterConfig.MeterEventService, |
| LLMCostService: config.RouterConfig.LLMCostService, |
| MeterService: config.RouterConfig.MeterManageService, |
| StreamingConnector: config.RouterConfig.StreamingConnector, |
| PlanService: config.RouterConfig.Plan, |
| PlanAddonService: config.RouterConfig.PlanAddon, |
| PlanSubscriptionService: config.RouterConfig.PlanSubscriptionService, |
| StripeService: config.RouterConfig.AppStripe, |
| SubscriptionService: config.RouterConfig.SubscriptionService, |
| SubscriptionAddonService: config.RouterConfig.SubscriptionAddonService, |
| SubscriptionWorkflowService: config.RouterConfig.SubscriptionWorkflowService, |
| ChargeService: config.RouterConfig.ChargeService, |
| TaxCodeService: config.RouterConfig.TaxCodeService, |
| CostService: config.RouterConfig.CostService, |
| FeatureConnector: config.RouterConfig.FeatureConnector, |
| Middlewares: v3Middlewares, |
| PostAuthMiddlewares: config.PostAuthMiddlewares, |
| ResponseValidation: config.ResponseValidation, |
| FeatureGate: config.RouterConfig.FeatureGate, |
| }) |
| if err != nil { |
| return nil, fmt.Errorf("failed to create v3 API: %w", err) |
| } |
|
|
| var v3RegisterErr error |
| r.Group(func(r chi.Router) { |
| v3RegisterErr = v3API.RegisterRoutes(r) |
| }) |
| if v3RegisterErr != nil { |
| return nil, fmt.Errorf("failed to register v3 API routes: %w", v3RegisterErr) |
| } |
|
|
| r.Group(func(r chi.Router) { |
| |
| for _, mw := range hookMiddlewares { |
| r.Use(mw) |
| } |
|
|
| r.Use(config.ClientIPMiddleware) |
| r.Use(middleware.RequestID) |
| r.Use(func(h http.Handler) http.Handler { |
| return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| ctx := r.Context() |
| ctx = contextx.WithAttrs(ctx, server.GetRequestAttributes(r)) |
|
|
| h.ServeHTTP(w, r.WithContext(ctx)) |
| }) |
| }) |
| r.Use(server.NewRequestLoggerMiddleware(slog.Default().Handler())) |
| r.Use(middleware.Recoverer) |
| if config.RouterConfig.PortalCORSEnabled { |
| |
| r.Use(corsHandler(corsOptions{ |
| AllowedPaths: []string{"/api/v1/portal/meters"}, |
| Options: cors.Options{ |
| AllowOriginFunc: func(r *http.Request, origin string) bool { |
| return true |
| }, |
| AllowedMethods: []string{http.MethodGet, http.MethodOptions}, |
| AllowedHeaders: []string{"Accept", "Authorization", "Content-Type"}, |
| AllowCredentials: true, |
| MaxAge: 1728000, |
| }, |
| })) |
| } |
| r.Use(render.SetContentType(render.ContentTypeJSON)) |
| r.NotFound(func(w http.ResponseWriter, r *http.Request) { |
| models.NewStatusProblem(r.Context(), nil, http.StatusNotFound).Respond(w) |
| }) |
| r.MethodNotAllowed(func(w http.ResponseWriter, r *http.Request) { |
| models.NewStatusProblem(r.Context(), nil, http.StatusMethodNotAllowed).Respond(w) |
| }) |
|
|
| |
| r.Get("/api/swagger.json", func(w http.ResponseWriter, r *http.Request) { |
| render.JSON(w, r, swagger) |
| }) |
|
|
| |
| for _, routeHook := range config.RouterHooks.Routes { |
| routeHook(r) |
| } |
|
|
| middlewares := []api.MiddlewareFunc{ |
| authenticator.NewAuthenticator(config.RouterConfig.Portal, config.RouterConfig.ErrorHandler).NewAuthenticatorMiddlewareFunc(swagger), |
| oapimiddleware.OapiRequestValidatorWithOptions(swagger, &oapimiddleware.Options{ |
| ErrorHandler: func(w http.ResponseWriter, message string, statusCode int) { |
| models.NewStatusProblem(context.Background(), errors.New(message), statusCode).Respond(w) |
| }, |
| Options: openapi3filter.Options{ |
| |
| AuthenticationFunc: openapi3filter.NoopAuthenticationFunc, |
| SkipSettingDefaults: true, |
|
|
| |
| |
| ExcludeReadOnlyValidations: true, |
| }, |
| }), |
| } |
|
|
| postAuthMiddlewares := lo.Map(config.PostAuthMiddlewares, func(mwf server.MiddlewareFunc, _ int) api.MiddlewareFunc { |
| return api.MiddlewareFunc(mwf) |
| }) |
|
|
| middlewares = append(middlewares, postAuthMiddlewares...) |
|
|
| |
| _ = api.HandlerWithOptions(impl, api.ChiServerOptions{ |
| BaseRouter: r, |
| Middlewares: middlewares, |
| ErrorHandlerFunc: func(w http.ResponseWriter, r *http.Request, err error) { |
| config.RouterConfig.ErrorHandler.HandleContext(r.Context(), err) |
| errorHandlerReply(w, r, err) |
| }, |
| }) |
| }) |
|
|
| return &Server{ |
| Router: r, |
| }, nil |
| } |
|
|
| |
| type middlewareCollector struct { |
| middlewares []server.MiddlewareFunc |
| } |
|
|
| func (c *middlewareCollector) Use(middlewares ...func(http.Handler) http.Handler) { |
| for _, mw := range middlewares { |
| c.middlewares = append(c.middlewares, server.MiddlewareFunc(mw)) |
| } |
| } |
|
|
| |
| func collectMiddlewareHooks(hooks []MiddlewareHook) []server.MiddlewareFunc { |
| c := &middlewareCollector{} |
| for _, hook := range hooks { |
| hook(c) |
| } |
| return c.middlewares |
| } |
|
|
| |
| func errorHandlerReply(w http.ResponseWriter, r *http.Request, err error) { |
| switch e := err.(type) { |
| case *api.UnescapedCookieParamError: |
| err := fmt.Errorf("unescaped cookie param %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| case *api.UnmarshalingParamError: |
| err := fmt.Errorf("unmarshaling param %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| case *api.RequiredParamError: |
| err := fmt.Errorf("required param missing %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| case *api.RequiredHeaderError: |
| err := fmt.Errorf("required header missing %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| case *api.InvalidParamFormatError: |
| err := fmt.Errorf("invalid param format %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| case *api.TooManyValuesForParamError: |
| err := fmt.Errorf("too many values for param %s: %w", e.ParamName, err) |
| models.NewStatusProblem(r.Context(), err, http.StatusBadRequest).Respond(w) |
| default: |
| err := fmt.Errorf("unhandled server error: %w", err) |
| models.NewStatusProblem(r.Context(), err, http.StatusInternalServerError).Respond(w) |
| } |
| } |
|
|