oto/services/core/internal/httpserver/runner_handlers.go

452 lines
13 KiB
Go

package httpserver
import (
"encoding/json"
"fmt"
"net"
"net/http"
"net/url"
"os"
"strings"
"github.com/toki/oto/services/core/internal/runnerregistry"
otopb "github.com/toki/oto/services/core/oto"
)
func handleHealthz(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
w.WriteHeader(http.StatusOK)
w.Write([]byte("OK"))
}
func handleReadyz(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
w.WriteHeader(http.StatusOK)
w.Write([]byte("OK"))
}
func handleRunnerRegister(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
var request otopb.RegisterRunnerRequest
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
writeRunnerRegisterResponse(w, http.StatusBadRequest, &otopb.RegisterRunnerResponse{
Accepted: false,
RejectReason: "invalid registration request",
Error: &otopb.ProtocolError{
Code: "invalid_registration_request",
Message: "invalid registration request",
},
})
return
}
response := registry.Register(&request)
writeRunnerRegisterResponse(w, http.StatusOK, response)
}
}
func handleRunnerHeartbeat(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
runnerID := r.PathValue("id")
if runnerID == "" {
parts := strings.Split(r.URL.Path, "/")
if len(parts) >= 5 {
runnerID = parts[4]
}
}
if runnerID == "" {
writeHeartbeatResponse(w, http.StatusBadRequest, &otopb.HeartbeatResponse{
Success: false,
ErrorMessage: "missing runner id",
})
return
}
var req otopb.HeartbeatRequest
if r.ContentLength > 0 || r.Body != nil {
if err := json.NewDecoder(r.Body).Decode(&req); err != nil && err.Error() != "EOF" {
writeHeartbeatResponse(w, http.StatusBadRequest, &otopb.HeartbeatResponse{
Success: false,
ErrorMessage: "invalid heartbeat request",
})
return
}
}
if req.RunnerId == "" {
req.RunnerId = runnerID
} else if req.RunnerId != runnerID {
writeHeartbeatResponse(w, http.StatusBadRequest, &otopb.HeartbeatResponse{
Success: false,
ErrorMessage: "runner id mismatch between path and body",
})
return
}
response := registry.Heartbeat(req.RunnerId, req.Status)
if !response.Success {
status := http.StatusNotFound
if response.ErrorMessage != "unknown runner" {
status = http.StatusBadRequest
}
writeHeartbeatResponse(w, status, response)
return
}
writeHeartbeatResponse(w, http.StatusOK, response)
}
}
func handleRunnerDisconnect(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
runnerID := r.PathValue("id")
if runnerID == "" {
parts := strings.Split(r.URL.Path, "/")
if len(parts) >= 5 {
runnerID = parts[4]
}
}
if runnerID == "" {
writeResponse(w, http.StatusBadRequest, runnerDisconnectToJSON(false, "missing runner id"))
return
}
success := registry.Disconnect(runnerID)
if !success {
writeResponse(w, http.StatusNotFound, runnerDisconnectToJSON(false, "runner not found"))
return
}
writeResponse(w, http.StatusOK, runnerDisconnectToJSON(true, ""))
}
}
func handleGetRunner(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
runnerID := r.PathValue("id")
if runnerID == "" {
parts := strings.Split(r.URL.Path, "/")
if len(parts) >= 5 {
runnerID = parts[4]
}
}
if runnerID == "" {
http.Error(w, "missing runner id", http.StatusBadRequest)
return
}
record, ok := registry.Snapshot(runnerID)
if !ok {
http.Error(w, "runner not found", http.StatusNotFound)
return
}
writeResponse(w, http.StatusOK, runnerRecordToJSON(record))
}
}
func handleUpdateRunner(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPatch {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
runnerID := r.PathValue("id")
if runnerID == "" {
parts := strings.Split(r.URL.Path, "/")
if len(parts) >= 5 {
runnerID = parts[4]
}
}
if runnerID == "" {
writeResponse(w, http.StatusBadRequest, errorToJSON("missing runner id"))
return
}
var request struct {
Alias string `json:"alias"`
}
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
writeResponse(w, http.StatusBadRequest, errorToJSON("invalid runner update request"))
return
}
alias := strings.TrimSpace(request.Alias)
if alias == "" {
writeResponse(w, http.StatusBadRequest, errorToJSON("missing alias"))
return
}
record, ok := registry.UpdateAlias(runnerID, alias)
if !ok {
writeResponse(w, http.StatusNotFound, errorToJSON("runner not found"))
return
}
writeResponse(w, http.StatusOK, runnerRecordToJSON(record))
}
}
func handleRunnerResource(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
handleGetRunner(registry)(w, r)
case http.MethodPatch:
handleUpdateRunner(registry)(w, r)
default:
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
}
}
}
func shellEscape(s string) string {
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
}
func powershellEscape(s string) string {
return "'" + strings.ReplaceAll(s, "'", "''") + "'"
}
func handleRunnerBootstrapCommand(registry *runnerregistry.Registry) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
var request struct {
EnrollmentToken string `json:"enrollment_token"`
Token string `json:"token"`
Target string `json:"target"`
}
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
writeResponse(w, http.StatusBadRequest, errorToJSON("invalid bootstrap command request"))
return
}
token := strings.TrimSpace(request.Token)
if token == "" {
token = strings.TrimSpace(request.EnrollmentToken)
}
if token == "" {
writeResponse(w, http.StatusBadRequest, errorToJSON("missing token"))
return
}
target := strings.ToLower(strings.TrimSpace(request.Target))
if target != "" && target != "linux" && target != "macos" && target != "windows" {
writeResponse(w, http.StatusBadRequest, errorToJSON(fmt.Sprintf("unsupported bootstrap target: %s", target)))
return
}
defaults, err := bootstrapDefaultsForRequest(r, true)
if err != nil {
writeResponse(w, http.StatusBadRequest, errorToJSON(err.Error()))
return
}
serverURL := defaults.ServerURL
var bootstrapCmd string
if target == "windows" {
escapedScriptURL := powershellEscape(serverURL + "/bootstrap/oto-agent.ps1")
escapedToken := powershellEscape(token)
bootstrapCmd = fmt.Sprintf(
"powershell -ExecutionPolicy Bypass -Command \"& ([scriptblock]::Create((irm %s))) -- --token %s\"",
escapedScriptURL, escapedToken,
)
} else {
escapedScriptURL := shellEscape(serverURL + "/bootstrap/oto-agent.sh")
escapedToken := shellEscape(token)
bootstrapCmd = fmt.Sprintf(
"curl -fsSL %s | bash -s -- --token %s",
escapedScriptURL, escapedToken,
)
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_ = json.NewEncoder(w).Encode(&otopb.BootstrapCommandResponse{
BootstrapCommand: bootstrapCmd,
})
}
}
type bootstrapDefaults struct {
ServerURL string
SocketURL string
ReleaseBaseURL string
}
func bootstrapDefaultsForRequest(r *http.Request, requireRelease bool) (bootstrapDefaults, error) {
scheme := "http"
if r.TLS != nil {
scheme = "https"
}
serverURL := scheme + "://" + r.Host
u, err := url.Parse(serverURL)
if err != nil || u.Host == "" || u.Scheme == "" {
return bootstrapDefaults{}, fmt.Errorf("invalid server URL")
}
// Host validation to prevent command injection and host header spoofing.
for _, char := range u.Host {
if !((char >= 'a' && char <= 'z') || (char >= 'A' && char <= 'Z') || (char >= '0' && char <= '9') ||
char == '.' || char == '-' || char == ':' || char == '[' || char == ']') {
return bootstrapDefaults{}, fmt.Errorf("invalid characters in server Host")
}
}
releaseBaseURL := strings.TrimSpace(os.Getenv("OTO_RUNNER_RELEASE_BASE_URL"))
if releaseBaseURL == "" {
if scheme == "https" {
releaseBaseURL = serverURL + "/releases"
} else if requireRelease {
return bootstrapDefaults{}, fmt.Errorf("release base URL must use HTTPS (set OTO_RUNNER_RELEASE_BASE_URL environment variable)")
}
} else if !strings.HasPrefix(releaseBaseURL, "https://") {
if requireRelease {
return bootstrapDefaults{}, fmt.Errorf("OTO_RUNNER_RELEASE_BASE_URL must use https:// scheme")
}
releaseBaseURL = ""
}
socketURL, err := runnerSocketPublicURL(serverURL)
if err != nil {
if requireRelease {
return bootstrapDefaults{}, err
}
socketURL = ""
}
return bootstrapDefaults{
ServerURL: serverURL,
SocketURL: socketURL,
ReleaseBaseURL: releaseBaseURL,
}, nil
}
func runnerSocketPublicURL(serverURL string) (string, error) {
if configured := strings.TrimSpace(os.Getenv("OTO_RUNNER_SOCKET_PUBLIC_URL")); configured != "" {
if !strings.Contains(configured, "://") {
configured = "tcp://" + configured
}
u, err := url.Parse(configured)
if err != nil || u.Hostname() == "" {
return "", fmt.Errorf("invalid OTO_RUNNER_SOCKET_PUBLIC_URL")
}
return configured, nil
}
u, err := url.Parse(serverURL)
if err != nil || u.Hostname() == "" {
return "", fmt.Errorf("invalid server URL")
}
return "tcp://" + net.JoinHostPort(u.Hostname(), "18080"), nil
}
func handleServeBootstrapScript(provider bootstrapScriptProvider) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
content, err := provider.bootstrapScript()
if err != nil {
http.Error(w, "Bootstrap script not available: "+err.Error(), http.StatusInternalServerError)
return
}
defaults, _ := bootstrapDefaultsForRequest(r, false)
content = renderShellBootstrapDefaults(content, defaults)
w.Header().Set("Content-Type", "application/x-sh")
w.WriteHeader(http.StatusOK)
_, _ = w.Write(content)
}
}
func handleServeBootstrapPs1(provider bootstrapPs1Provider) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
return
}
content, err := provider.bootstrapPs1()
if err != nil {
http.Error(w, "Bootstrap PS1 script not available: "+err.Error(), http.StatusInternalServerError)
return
}
defaults, _ := bootstrapDefaultsForRequest(r, false)
content = renderPowerShellBootstrapDefaults(content, defaults)
w.Header().Set("Content-Type", "application/x-powershell")
w.WriteHeader(http.StatusOK)
_, _ = w.Write(content)
}
}
func renderShellBootstrapDefaults(content []byte, defaults bootstrapDefaults) []byte {
rendered := string(content)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_SERVER_URL__", shellSingleQuoteContent(defaults.ServerURL), 1)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_SOCKET_URL__", shellSingleQuoteContent(defaults.SocketURL), 1)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_RELEASE_BASE_URL__", shellSingleQuoteContent(defaults.ReleaseBaseURL), 1)
return []byte(rendered)
}
func renderPowerShellBootstrapDefaults(content []byte, defaults bootstrapDefaults) []byte {
rendered := string(content)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_SERVER_URL__", powershellSingleQuoteContent(defaults.ServerURL), 1)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_SOCKET_URL__", powershellSingleQuoteContent(defaults.SocketURL), 1)
rendered = strings.Replace(rendered, "__OTO_BOOTSTRAP_DEFAULT_RELEASE_BASE_URL__", powershellSingleQuoteContent(defaults.ReleaseBaseURL), 1)
return []byte(rendered)
}
func shellSingleQuoteContent(s string) string {
return strings.ReplaceAll(s, "'", "'\\''")
}
func powershellSingleQuoteContent(s string) string {
return strings.ReplaceAll(s, "'", "''")
}
func writeRunnerRegisterResponse(w http.ResponseWriter, status int, response *otopb.RegisterRunnerResponse) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(response)
}
func writeHeartbeatResponse(w http.ResponseWriter, status int, response *otopb.HeartbeatResponse) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(response)
}