From af221cf3c15a96bf18ac69999129d86bae741c37 Mon Sep 17 00:00:00 2001 From: Nikolay Edigaryev Date: Mon, 6 Oct 2025 18:04:47 +0200 Subject: [PATCH] Support for prefixed Orchard Controller API URLs (#355) * Support for prefixed Orchard Controller API URLs * Fix Swagger UI * Remove spurious "fmt" import * Use url.URL in order to correctly calculate API path for Swagger UI --- internal/command/controller/run.go | 8 ++++++++ internal/command/dev/dev.go | 15 +++++++++++--- internal/controller/api.go | 24 +++++++++++++++++----- internal/controller/controller.go | 14 ++++++++++--- internal/controller/option.go | 6 ++++++ internal/netconstants/netconstants_test.go | 22 ++++++++++++++++++++ pkg/client/client.go | 12 ++++++----- 7 files changed, 85 insertions(+), 16 deletions(-) create mode 100644 internal/netconstants/netconstants_test.go diff --git a/internal/command/controller/run.go b/internal/command/controller/run.go index 3313e6a..66fb7b4 100644 --- a/internal/command/controller/run.go +++ b/internal/command/controller/run.go @@ -24,6 +24,7 @@ import ( var ErrRunFailed = errors.New("failed to run controller") var address string +var apiPrefix string var addressSSH string var addressPprof string var debug bool @@ -50,6 +51,9 @@ func newRunCommand() *cobra.Command { cmd.Flags().StringVarP(&address, "listen", "l", fmt.Sprintf(":%s", port), "address to listen on") + cmd.Flags().StringVar(&apiPrefix, "api-prefix", "", + "prefix to prepend to all Orchard Controller API endpoints; useful when exposing Orchard Controller "+ + "behind an HTTP proxy together with other services") cmd.Flags().StringVar(&addressSSH, "listen-ssh", "", "address for the built-in SSH server to listen on (e.g. \":6122\")") cmd.Flags().StringVar(&addressPprof, "listen-pprof", "", @@ -144,6 +148,10 @@ func runController(cmd *cobra.Command, args []string) (err error) { controller.WithLogger(logger), } + if apiPrefix != "" { + controllerOpts = append(controllerOpts, controller.WithAPIPrefix(apiPrefix)) + } + var controllerCert tls.Certificate if !noTLS { diff --git a/internal/command/dev/dev.go b/internal/command/dev/dev.go index 662d83d..0771309 100644 --- a/internal/command/dev/dev.go +++ b/internal/command/dev/dev.go @@ -5,6 +5,10 @@ package dev import ( "errors" "fmt" + "os" + "path" + "path/filepath" + "github.com/cirruslabs/orchard/internal/config" "github.com/cirruslabs/orchard/internal/controller" "github.com/cirruslabs/orchard/internal/netconstants" @@ -13,14 +17,12 @@ import ( v1 "github.com/cirruslabs/orchard/pkg/resource/v1" "github.com/spf13/cobra" "go.uber.org/zap" - "os" - "path" - "path/filepath" ) var ErrFailed = errors.New("failed to run development controller and worker") var devDataDirPath string +var apiPrefix string var stringToStringResources map[string]string var experimentalRPCV2 bool @@ -33,6 +35,9 @@ func NewCommand() *cobra.Command { command.Flags().StringVarP(&devDataDirPath, "data-dir", "d", ".dev-data", "path to persist data between runs") + command.Flags().StringVar(&apiPrefix, "api-prefix", "", + "prefix to prepend to all Orchard Controller API endpoints; useful when exposing Orchard Controller "+ + "behind an HTTP proxy together with other services") command.Flags().StringToStringVar(&stringToStringResources, "resources", map[string]string{}, "resources that the development worker will provide") command.Flags().BoolVar(&experimentalRPCV2, "experimental-rpc-v2", false, @@ -58,6 +63,10 @@ func runDev(cmd *cobra.Command, args []string) error { var additionalControllerOpts []controller.Option + if apiPrefix != "" { + additionalControllerOpts = append(additionalControllerOpts, controller.WithAPIPrefix(apiPrefix)) + } + if experimentalRPCV2 { additionalControllerOpts = append(additionalControllerOpts, controller.WithExperimentalRPCV2()) } diff --git a/internal/controller/api.go b/internal/controller/api.go index 023b1fa..99070e8 100644 --- a/internal/controller/api.go +++ b/internal/controller/api.go @@ -5,6 +5,7 @@ import ( "crypto/subtle" "errors" "net/http" + "net/url" "strings" "github.com/cirruslabs/orchard/api" @@ -28,7 +29,15 @@ var ErrUnauthorized = errors.New("unauthorized") func (controller *Controller) initAPI() *gin.Engine { ginEngine := gin.New() - ginEngine.Use( + var group *gin.RouterGroup + + if controller.apiPrefix != "" { + group = ginEngine.Group(controller.apiPrefix) + } else { + group = ginEngine.Group("/") + } + + group.Use( ginzap.Ginzap(controller.logger.Desugar(), "", true), ginzap.RecoveryWithZap(controller.logger.Desugar(), true), ) @@ -36,10 +45,10 @@ func (controller *Controller) initAPI() *gin.Engine { // expose metrics monitor := ginmetrics.GetMonitor() monitor.SetMetricPath("/metrics") - monitor.Use(ginEngine) + monitor.Use(group) // v1 API - v1 := ginEngine.Group("/v1") + v1 := group.Group("/v1") // Auth v1.Use(controller.authenticateMiddleware) @@ -48,9 +57,14 @@ func (controller *Controller) initAPI() *gin.Engine { // to check that the API is working v1.GET("/", func(c *gin.Context) { if controller.enableSwaggerDocs { + apiURL := &url.URL{ + Path: "/", + } + apiURL = apiURL.JoinPath(controller.apiPrefix, "v1") + middleware.SwaggerUI(middleware.SwaggerUIOpts{ - Path: "/v1", - SpecURL: "/v1/openapi.yaml", + Path: apiURL.Path, + SpecURL: apiURL.JoinPath("openapi.yaml").Path, }, nil).ServeHTTP(c.Writer, c.Request) } else { c.Status(http.StatusOK) diff --git a/internal/controller/controller.go b/internal/controller/controller.go index 3d1bc7b..cb26602 100644 --- a/internal/controller/controller.go +++ b/internal/controller/controller.go @@ -7,6 +7,7 @@ import ( "fmt" "net" "net/http" + "net/url" "os" "strings" "time" @@ -46,6 +47,7 @@ const ( type Controller struct { dataDir *DataDir listenAddr string + apiPrefix string tlsConfig *tls.Config listener net.Listener httpServer *http.Server @@ -262,11 +264,17 @@ func (controller *Controller) Run(ctx context.Context) error { func (controller *Controller) Address() string { hostPort := strings.ReplaceAll(controller.listener.Addr().String(), "[::]", "127.0.0.1") - if controller.tlsConfig != nil { - return fmt.Sprintf("https://%s", hostPort) + url := url.URL{ + Scheme: "http", + Host: hostPort, + Path: controller.apiPrefix, } - return fmt.Sprintf("http://%s", hostPort) + if controller.tlsConfig != nil { + url.Scheme = "https" + } + + return url.String() } func (controller *Controller) SSHAddress() (string, bool) { diff --git a/internal/controller/option.go b/internal/controller/option.go index e30384a..a005d68 100644 --- a/internal/controller/option.go +++ b/internal/controller/option.go @@ -22,6 +22,12 @@ func WithListenAddr(listenAddr string) Option { } } +func WithAPIPrefix(apiPrefix string) Option { + return func(c *Controller) { + c.apiPrefix = apiPrefix + } +} + func WithTLSConfig(tlsConfig *tls.Config) Option { return func(controller *Controller) { controller.tlsConfig = tlsConfig diff --git a/internal/netconstants/netconstants_test.go b/internal/netconstants/netconstants_test.go new file mode 100644 index 0000000..c827ca6 --- /dev/null +++ b/internal/netconstants/netconstants_test.go @@ -0,0 +1,22 @@ +package netconstants_test + +import ( + "testing" + + "github.com/cirruslabs/orchard/internal/netconstants" + "github.com/stretchr/testify/require" +) + +func TestNormalizeAddress(t *testing.T) { + // Default port + url, err := netconstants.NormalizeAddress("subdomain.example.com/some/prefix") + require.NoError(t, err) + require.Equal(t, "subdomain.example.com:6120", url.Host) + require.Equal(t, "/some/prefix", url.Path) + + // Custom port + url, err = netconstants.NormalizeAddress("subdomain.example.com:443/some/prefix") + require.NoError(t, err) + require.Equal(t, "subdomain.example.com:443", url.Host) + require.Equal(t, "/some/prefix", url.Path) +} diff --git a/pkg/client/client.go b/pkg/client/client.go index 6406d68..f7f6315 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -9,6 +9,12 @@ import ( "encoding/json" "errors" "fmt" + "io" + "net" + "net/http" + "net/url" + "time" + "github.com/cirruslabs/orchard/internal/config" "github.com/cirruslabs/orchard/internal/version" "github.com/cirruslabs/orchard/rpc" @@ -16,11 +22,6 @@ import ( "google.golang.org/grpc/credentials" "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/metadata" - "io" - "net" - "net/http" - "net/url" - "time" ) var ( @@ -298,6 +299,7 @@ func (client *Client) formatPath(path string) *url.URL { Scheme: client.baseURL.Scheme, User: client.baseURL.User, Host: client.baseURL.Host, + Path: client.baseURL.Path, } return endpointURL.JoinPath("v1", path)