forked from aegis/pyserveX
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bd0b381195 | ||
|
|
81ac5c4d29 | ||
|
|
881028c1e6 | ||
|
|
8f5b9a5cd1 | ||
|
|
c04ab283a6 | ||
|
|
d03ade18c5 | ||
|
|
129785706c | ||
|
|
3b59994fc9 | ||
|
|
7662a7924a | ||
|
|
cec6e927a7 | ||
|
|
80544d5b95 | ||
|
|
b4f63c6804 | ||
|
|
59d6ae2fd2 | ||
|
|
edaccb59bb | ||
|
|
3454801be7 |
+4
-1
@@ -27,4 +27,7 @@ build/
|
||||
.idea/
|
||||
.vscode/
|
||||
*.swp
|
||||
*.swo
|
||||
*.swo
|
||||
|
||||
# Go binaries
|
||||
go/bin
|
||||
@@ -1,145 +1,97 @@
|
||||
# PyServe
|
||||
|
||||
PyServe is a modern, async HTTP server written in Python. Originally created for educational purposes, it has evolved into a powerful tool for rapid prototyping and serving web applications with unique features like AI-generated content.
|
||||
Python application orchestrator and HTTP server. Runs multiple ASGI/WSGI applications through a single entry point with process isolation, health monitoring, and auto-restart.
|
||||
|
||||
<img src="https://raw.githubusercontent.com/ShiftyX1/PyServe/refs/heads/master/images/logo.png" alt="isolated" width="150"/>
|
||||
<img src="https://raw.githubusercontent.com/ShiftyX1/PyServe/refs/heads/master/images/logo.png" alt="PyServe Logo" width="150"/>
|
||||
|
||||
[More on web page](https://pyserve.org/)
|
||||
Website: [pyserve.org](https://pyserve.org) · Documentation: [docs.pyserve.org](https://docs.pyserve.org)
|
||||
|
||||
## Project Overview
|
||||
## Overview
|
||||
|
||||
PyServe v0.6.0 introduces a completely refactored architecture with modern async/await syntax and new exciting features like **Vibe-Serving** - AI-powered dynamic content generation.
|
||||
PyServe manages multiple Python web applications (FastAPI, Flask, Django, etc.) as isolated subprocesses behind a single gateway. Each app runs on its own port with independent lifecycle, health checks, and automatic restarts on failure.
|
||||
|
||||
### Key Features:
|
||||
```
|
||||
PyServe Gateway (:8000)
|
||||
│
|
||||
┌────────────────┼────────────────┐
|
||||
▼ ▼ ▼
|
||||
FastAPI Flask Starlette
|
||||
:9001 :9002 :9003
|
||||
/api/* /admin/* /ws/*
|
||||
```
|
||||
|
||||
- **Async HTTP Server** - Built with Python's asyncio for high performance
|
||||
- **Advanced Configuration System V2** - Powerful extensible configuration with full backward compatibility
|
||||
- **Regex Routing & SPA Support** - nginx-style routing patterns with Single Page Application fallback
|
||||
- **Static File Serving** - Efficient serving with correct MIME types
|
||||
- **Template System** - Dynamic content generation
|
||||
- **Vibe-Serving Mode** - AI-generated content using language models (OpenAI, Claude, etc.)
|
||||
- **Reverse Proxy** - Forward requests to backend services with advanced routing
|
||||
- **SSL/HTTPS Support** - Secure connections with certificate configuration
|
||||
- **Modular Extensions** - Plugin-like architecture for security, caching, monitoring
|
||||
- **Beautiful Logging** - Colored terminal output with file rotation
|
||||
- **Error Handling** - Styled error pages and graceful fallbacks
|
||||
- **CLI Interface** - Command-line interface for easy deployment and configuration
|
||||
|
||||
## Getting Started
|
||||
|
||||
### Prerequisites
|
||||
|
||||
- Python 3.12 or higher
|
||||
- Poetry (recommended) or pip
|
||||
|
||||
### Installation
|
||||
|
||||
#### Via Poetry (рекомендуется)
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ShiftyX1/PyServe.git
|
||||
cd PyServe
|
||||
make init # Initialize project
|
||||
make init
|
||||
```
|
||||
|
||||
#### Или установка пакета
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# local install
|
||||
make install-package
|
||||
```yaml
|
||||
# config.yaml
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8000
|
||||
|
||||
# after installing project you can use command pyserve
|
||||
pyserve --help
|
||||
extensions:
|
||||
- type: process_orchestration
|
||||
config:
|
||||
apps:
|
||||
- name: api
|
||||
path: /api
|
||||
app_path: myapp.api:app
|
||||
|
||||
- name: admin
|
||||
path: /admin
|
||||
app_path: myapp.admin:app
|
||||
```
|
||||
|
||||
### Running the Server
|
||||
|
||||
#### Using Makefile (recommended)
|
||||
|
||||
```bash
|
||||
# start in development mode
|
||||
make run
|
||||
|
||||
# start in production mode
|
||||
make run-prod
|
||||
|
||||
# show all available commands
|
||||
make help
|
||||
pyserve -c config.yaml
|
||||
```
|
||||
|
||||
#### Using CLI directly
|
||||
Requests to `/api/*` are proxied to the api process, `/admin/*` to admin.
|
||||
|
||||
## Process Orchestration
|
||||
|
||||
The main case of using PyServe is orchestration of python web applications. Each application runs as a separate uvicorn process on a dynamically or manually allocated port (9000-9999 by default). PyServe proxies requests to the appropriate process based on URL path.
|
||||
|
||||
For each application you can configure the number of workers, environment variables, health check endpoint path, and auto-restart parameters. If a process crashes or stops responding to health checks, PyServe automatically restarts it with exponential backoff.
|
||||
|
||||
WSGI applications (Flask, Django) are supported through automatic wrapping — just specify `app_type: wsgi`.
|
||||
|
||||
## In-Process Mounting
|
||||
|
||||
For simpler cases when process isolation is not needed, applications can be mounted directly into the PyServe process via the `asgi` extension. This is lighter and faster, but all applications share one process.
|
||||
|
||||
## Static Files & Routing
|
||||
|
||||
PyServe can serve static files with nginx-like routing: regex patterns, SPA fallback for frontend applications, custom caching headers. Routes are processed in priority order — exact match, then regex, then default.
|
||||
|
||||
## Reverse Proxy
|
||||
|
||||
Requests can be proxied to external backends. Useful for integration with legacy services or microservices in other languages.
|
||||
|
||||
## CLI
|
||||
|
||||
```bash
|
||||
# after installing package
|
||||
pyserve
|
||||
|
||||
# or with Poetry
|
||||
poetry run pyserve
|
||||
|
||||
# or legacy (for backward compatibility)
|
||||
python run.py
|
||||
```
|
||||
|
||||
#### CLI options
|
||||
|
||||
```bash
|
||||
# help
|
||||
pyserve --help
|
||||
|
||||
# path to config
|
||||
pyserve -c /path/to/config.yaml
|
||||
|
||||
# rewrite host and port
|
||||
pyserve -c config.yaml
|
||||
pyserve --host 0.0.0.0 --port 9000
|
||||
|
||||
# debug mode
|
||||
pyserve --debug
|
||||
|
||||
# show version
|
||||
pyserve --version
|
||||
```
|
||||
|
||||
## Development
|
||||
|
||||
### Makefile Commands
|
||||
|
||||
```bash
|
||||
make help # Show help for commands
|
||||
make install # Install dependencies
|
||||
make dev-install # Install development dependencies
|
||||
make build # Build the package
|
||||
make test # Run tests
|
||||
make test-cov # Tests with code coverage
|
||||
make lint # Check code with linters
|
||||
make format # Format code
|
||||
make clean # Clean up temporary files
|
||||
make version # Show version
|
||||
make publish-test # Publish to Test PyPI
|
||||
make publish # Publish to PyPI
|
||||
```
|
||||
|
||||
### Project Structure
|
||||
|
||||
```
|
||||
pyserveX/
|
||||
├── pyserve/ # Main package
|
||||
│ ├── __init__.py
|
||||
│ ├── cli.py # CLI interface
|
||||
│ ├── server.py # Main server module
|
||||
│ ├── config.py # Configuration system
|
||||
│ ├── routing.py # Routing
|
||||
│ ├── extensions.py # Extensions
|
||||
│ └── logging_utils.py
|
||||
├── tests/ # Tests
|
||||
├── static/ # Static files
|
||||
├── templates/ # Templates
|
||||
├── logs/ # Logs
|
||||
├── Makefile # Automation tasks
|
||||
├── pyproject.toml # Project configuration
|
||||
├── config.yaml # Server configuration
|
||||
└── run.py # Entry point (backward compatibility)
|
||||
make test # run tests
|
||||
make lint # linting
|
||||
make format # formatting
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
This project is distributed under the MIT license.
|
||||
[MIT License](./LICENSE)
|
||||
|
||||
@@ -1,32 +0,0 @@
|
||||
# PyServe configuration for serving documentation
|
||||
# Usage: pyserve -c config.docs.yaml
|
||||
|
||||
http:
|
||||
static_dir: ./docs
|
||||
templates_dir: ./templates
|
||||
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8000
|
||||
backlog: 5
|
||||
default_root: false
|
||||
|
||||
ssl:
|
||||
enabled: false
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
|
||||
extensions:
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
"~*\\.(css)$":
|
||||
root: "./docs"
|
||||
cache_control: "public, max-age=3600"
|
||||
|
||||
"__default__":
|
||||
root: "./docs"
|
||||
index_file: "index.html"
|
||||
cache_control: "no-cache"
|
||||
@@ -0,0 +1,113 @@
|
||||
# PyServe Process Orchestration Example
|
||||
#
|
||||
# This configuration demonstrates running multiple ASGI/WSGI applications
|
||||
# as isolated processes with automatic health monitoring and restart.
|
||||
|
||||
server:
|
||||
host: "0.0.0.0"
|
||||
port: 8000
|
||||
backlog: 2048
|
||||
proxy_timeout: 60.0
|
||||
|
||||
logging:
|
||||
level: DEBUG
|
||||
console_output: true
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
|
||||
extensions:
|
||||
# Process Orchestration - runs each app in its own process
|
||||
- type: process_orchestration
|
||||
config:
|
||||
# Port range for worker processes
|
||||
port_range: [9000, 9999]
|
||||
|
||||
# Enable health monitoring
|
||||
health_check_enabled: true
|
||||
|
||||
# Proxy timeout for requests
|
||||
proxy_timeout: 60.0
|
||||
|
||||
apps:
|
||||
# FastAPI application
|
||||
- name: api
|
||||
path: /api
|
||||
app_path: examples.apps.fastapi_app:app
|
||||
module_path: "."
|
||||
workers: 2
|
||||
health_check_path: /health
|
||||
health_check_interval: 10.0
|
||||
health_check_timeout: 5.0
|
||||
health_check_retries: 3
|
||||
max_restart_count: 5
|
||||
restart_delay: 1.0
|
||||
shutdown_timeout: 30.0
|
||||
strip_path: true
|
||||
env:
|
||||
APP_ENV: "production"
|
||||
DEBUG: "false"
|
||||
|
||||
# Flask application (WSGI wrapped to ASGI)
|
||||
- name: admin
|
||||
path: /admin
|
||||
app_path: examples.apps.flask_app:app
|
||||
app_type: wsgi
|
||||
module_path: "."
|
||||
workers: 1
|
||||
health_check_path: /health
|
||||
strip_path: true
|
||||
|
||||
# Starlette application
|
||||
- name: web
|
||||
path: /web
|
||||
app_path: examples.apps.starlette_app:app
|
||||
module_path: "."
|
||||
workers: 2
|
||||
health_check_path: /health
|
||||
strip_path: true
|
||||
|
||||
# Custom ASGI application
|
||||
- name: custom
|
||||
path: /custom
|
||||
app_path: examples.apps.custom_asgi:app
|
||||
module_path: "."
|
||||
workers: 1
|
||||
health_check_path: /health
|
||||
strip_path: true
|
||||
|
||||
# Routing for static files and reverse proxy
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
# Static files
|
||||
"^/static/.*":
|
||||
type: static
|
||||
root: "./static"
|
||||
strip_prefix: "/static"
|
||||
|
||||
# Documentation
|
||||
"^/docs/?.*":
|
||||
type: static
|
||||
root: "./docs"
|
||||
strip_prefix: "/docs"
|
||||
|
||||
# External API proxy
|
||||
"^/external/.*":
|
||||
type: proxy
|
||||
upstream: "https://api.example.com"
|
||||
strip_prefix: "/external"
|
||||
|
||||
# Security headers
|
||||
- type: security
|
||||
config:
|
||||
security_headers:
|
||||
X-Content-Type-Options: "nosniff"
|
||||
X-Frame-Options: "DENY"
|
||||
X-XSS-Protection: "1; mode=block"
|
||||
Strict-Transport-Security: "max-age=31536000; includeSubDomains"
|
||||
|
||||
# Monitoring
|
||||
- type: monitoring
|
||||
config:
|
||||
enable_metrics: true
|
||||
@@ -0,0 +1,34 @@
|
||||
# Multi-stage build for Konduktor
|
||||
FROM golang:1.23-alpine AS builder
|
||||
|
||||
RUN apk add --no-cache git make
|
||||
|
||||
WORKDIR /build
|
||||
|
||||
COPY go.mod go.sum* ./
|
||||
RUN go mod download
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN make build
|
||||
|
||||
FROM alpine:3.19
|
||||
|
||||
RUN apk add --no-cache ca-certificates tzdata
|
||||
|
||||
RUN adduser -D -g '' konduktor
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY --from=builder /build/bin/konduktor /usr/local/bin/
|
||||
COPY --from=builder /build/bin/konduktorctl /usr/local/bin/
|
||||
|
||||
RUN mkdir -p /app/static /app/templates /app/logs && \
|
||||
chown -R konduktor:konduktor /app
|
||||
|
||||
USER konduktor
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
ENTRYPOINT ["konduktor"]
|
||||
CMD ["-c", "/app/config.yaml"]
|
||||
+108
@@ -0,0 +1,108 @@
|
||||
# Konduktor Go Build
|
||||
# Makefile for building and testing Konduktor
|
||||
|
||||
.PHONY: all build build-konduktor build-konduktorctl test clean deps fmt lint run
|
||||
|
||||
# Build configuration
|
||||
VERSION ?= $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
|
||||
GIT_COMMIT ?= $(shell git rev-parse --short HEAD 2>/dev/null || echo "unknown")
|
||||
BUILD_TIME ?= $(shell date -u '+%Y-%m-%dT%H:%M:%SZ')
|
||||
LDFLAGS := -X main.Version=$(VERSION) -X main.GitCommit=$(GIT_COMMIT) -X main.BuildTime=$(BUILD_TIME)
|
||||
|
||||
# Output directories
|
||||
BIN_DIR := bin
|
||||
|
||||
all: deps build
|
||||
|
||||
# Download dependencies
|
||||
deps:
|
||||
@echo "==> Downloading dependencies..."
|
||||
go mod download
|
||||
go mod tidy
|
||||
|
||||
# Build all binaries
|
||||
build: build-konduktor build-konduktorctl
|
||||
|
||||
# Build konduktor server
|
||||
build-konduktor:
|
||||
@echo "==> Building konduktor..."
|
||||
@mkdir -p $(BIN_DIR)
|
||||
go build -ldflags "$(LDFLAGS)" -o $(BIN_DIR)/konduktor ./cmd/konduktor
|
||||
|
||||
# Build konduktorctl CLI
|
||||
build-konduktorctl:
|
||||
@echo "==> Building konduktorctl..."
|
||||
@mkdir -p $(BIN_DIR)
|
||||
go build -ldflags "$(LDFLAGS)" -o $(BIN_DIR)/konduktorctl ./cmd/konduktorctl
|
||||
|
||||
# Run tests
|
||||
test:
|
||||
@echo "==> Running tests..."
|
||||
go test -v -race -cover ./...
|
||||
|
||||
# Run tests with coverage report
|
||||
test-coverage:
|
||||
@echo "==> Running tests with coverage..."
|
||||
go test -v -race -coverprofile=coverage.out ./...
|
||||
go tool cover -html=coverage.out -o coverage.html
|
||||
@echo "Coverage report: coverage.html"
|
||||
|
||||
# Format code
|
||||
fmt:
|
||||
@echo "==> Formatting code..."
|
||||
go fmt ./...
|
||||
goimports -w .
|
||||
|
||||
# Lint code
|
||||
lint:
|
||||
@echo "==> Linting code..."
|
||||
golangci-lint run ./...
|
||||
|
||||
# Run the server (development)
|
||||
run: build-konduktor
|
||||
@echo "==> Running konduktor..."
|
||||
./$(BIN_DIR)/konduktor -c ../config.yaml
|
||||
|
||||
# Clean build artifacts
|
||||
clean:
|
||||
@echo "==> Cleaning..."
|
||||
rm -rf $(BIN_DIR)
|
||||
rm -f coverage.out coverage.html
|
||||
|
||||
# Install binaries to GOPATH/bin
|
||||
install: build
|
||||
@echo "==> Installing binaries..."
|
||||
cp $(BIN_DIR)/konduktor $(GOPATH)/bin/
|
||||
cp $(BIN_DIR)/konduktorctl $(GOPATH)/bin/
|
||||
|
||||
# Generate mocks (for testing)
|
||||
generate:
|
||||
@echo "==> Generating code..."
|
||||
go generate ./...
|
||||
|
||||
# Docker build
|
||||
docker-build:
|
||||
@echo "==> Building Docker image..."
|
||||
docker build -t konduktor:$(VERSION) .
|
||||
|
||||
# Show help
|
||||
help:
|
||||
@echo "Konduktor Build System"
|
||||
@echo ""
|
||||
@echo "Usage: make [target]"
|
||||
@echo ""
|
||||
@echo "Targets:"
|
||||
@echo " all Download deps and build all binaries"
|
||||
@echo " deps Download and tidy dependencies"
|
||||
@echo " build Build all binaries"
|
||||
@echo " build-konduktor Build the server binary"
|
||||
@echo " build-konduktorctl Build the CLI binary"
|
||||
@echo " test Run tests"
|
||||
@echo " test-coverage Run tests with coverage report"
|
||||
@echo " fmt Format code"
|
||||
@echo " lint Lint code"
|
||||
@echo " run Build and run the server"
|
||||
@echo " clean Clean build artifacts"
|
||||
@echo " install Install binaries to GOPATH/bin"
|
||||
@echo " docker-build Build Docker image"
|
||||
@echo " help Show this help"
|
||||
+149
@@ -0,0 +1,149 @@
|
||||
# Konduktor (Go)
|
||||
|
||||
High-performance HTTP web server with extensible routing and process orchestration. (Previously known as PyServe in Python)
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
go/
|
||||
├── cmd/
|
||||
│ ├── konduktor/ # Main server binary
|
||||
│ └── konduktorctl/ # CLI management tool
|
||||
├── internal/
|
||||
│ ├── config/ # Configuration management
|
||||
│ ├── logging/ # Structured logging
|
||||
│ ├── middleware/ # HTTP middleware
|
||||
│ ├── routing/ # HTTP routing
|
||||
│ ├── extensions/ # Extension system (TODO)
|
||||
│ └── process/ # Process management (TODO)
|
||||
├── pkg/ # Public packages (TODO)
|
||||
├── go.mod
|
||||
├── go.sum
|
||||
└── Makefile
|
||||
```
|
||||
|
||||
## Building
|
||||
|
||||
```bash
|
||||
cd go
|
||||
|
||||
# Download dependencies
|
||||
make deps
|
||||
|
||||
# Build all binaries
|
||||
make build
|
||||
|
||||
# Or build individually
|
||||
make build-konduktor
|
||||
make build-konduktorctl
|
||||
```
|
||||
|
||||
## Running
|
||||
|
||||
```bash
|
||||
# Run with default config
|
||||
./bin/konduktor
|
||||
|
||||
# Run with custom config
|
||||
./bin/konduktor -c ../config.yaml
|
||||
|
||||
# Run with flags
|
||||
./bin/konduktor --host 127.0.0.1 --port 3000 --debug
|
||||
```
|
||||
|
||||
## CLI Commands (konduktorctl)
|
||||
|
||||
```bash
|
||||
# Start services
|
||||
konduktorctl up
|
||||
|
||||
# Stop services
|
||||
konduktorctl down
|
||||
|
||||
# View status
|
||||
konduktorctl status
|
||||
|
||||
# View logs
|
||||
konduktorctl logs -f
|
||||
|
||||
# Health check
|
||||
konduktorctl health
|
||||
|
||||
# Scale services
|
||||
konduktorctl scale api=3
|
||||
|
||||
# Configuration management
|
||||
konduktorctl config show
|
||||
konduktorctl config validate
|
||||
|
||||
# Initialize new project
|
||||
konduktorctl init
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Uses the same YAML configuration format as the Python version:
|
||||
|
||||
```yaml
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
|
||||
http:
|
||||
static_dir: ./static
|
||||
templates_dir: ./templates
|
||||
|
||||
ssl:
|
||||
enabled: false
|
||||
cert_file: ./ssl/cert.pem
|
||||
key_file: ./ssl/key.pem
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
|
||||
extensions:
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
"=/health":
|
||||
return: "200 OK"
|
||||
```
|
||||
|
||||
## Development
|
||||
|
||||
```bash
|
||||
# Format code
|
||||
make fmt
|
||||
|
||||
# Run linter
|
||||
make lint
|
||||
|
||||
# Run tests
|
||||
make test
|
||||
|
||||
# Run with coverage
|
||||
make test-coverage
|
||||
```
|
||||
|
||||
## Migration from Python
|
||||
|
||||
This is a gradual rewrite of PyServe to Go. The project is now called **Konduktor**.
|
||||
|
||||
### Completed
|
||||
- [x] Basic project structure
|
||||
- [x] Configuration loading
|
||||
- [x] HTTP server with graceful shutdown
|
||||
- [x] Basic routing
|
||||
- [x] Middleware (access log, recovery, server header)
|
||||
- [x] CLI structure (konduktor, konduktorctl)
|
||||
|
||||
### TODO
|
||||
- [ ] Extension system
|
||||
- [x] Regex routing
|
||||
- [x] Reverse proxy
|
||||
- [ ] Process orchestration
|
||||
- [ ] ASGI/WSGI adapter support
|
||||
- [ ] WebSocket support
|
||||
- [ ] Hot reload
|
||||
- [ ] Metrics and monitoring
|
||||
@@ -0,0 +1,79 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/config"
|
||||
"github.com/konduktor/konduktor/internal/server"
|
||||
)
|
||||
|
||||
var (
|
||||
Version = "0.1.0"
|
||||
BuildTime = "unknown"
|
||||
GitCommit = "unknown"
|
||||
)
|
||||
|
||||
var (
|
||||
cfgFile string
|
||||
host string
|
||||
port int
|
||||
debug bool
|
||||
)
|
||||
|
||||
func main() {
|
||||
rootCmd := &cobra.Command{
|
||||
Use: "konduktor",
|
||||
Short: "Konduktor - HTTP web server",
|
||||
Long: `Konduktor is a high-performance HTTP web server with extensible routing and process orchestration.`,
|
||||
Version: fmt.Sprintf("%s (commit: %s, built: %s)", Version, GitCommit, BuildTime),
|
||||
RunE: runServer,
|
||||
}
|
||||
|
||||
rootCmd.Flags().StringVarP(&cfgFile, "config", "c", "config.yaml", "Path to configuration file")
|
||||
rootCmd.Flags().StringVar(&host, "host", "", "Host to bind the server to")
|
||||
rootCmd.Flags().IntVar(&port, "port", 0, "Port to bind the server to")
|
||||
rootCmd.Flags().BoolVar(&debug, "debug", false, "Enable debug mode")
|
||||
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func runServer(cmd *cobra.Command, args []string) error {
|
||||
cfg, err := config.Load(cfgFile)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
fmt.Printf("Configuration file %s not found, using defaults\n", cfgFile)
|
||||
cfg = config.Default()
|
||||
} else {
|
||||
return fmt.Errorf("configuration loading error: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if host != "" {
|
||||
cfg.Server.Host = host
|
||||
}
|
||||
if port != 0 {
|
||||
cfg.Server.Port = port
|
||||
}
|
||||
if debug {
|
||||
cfg.Logging.Level = "DEBUG"
|
||||
}
|
||||
|
||||
srv, err := server.New(cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("server creation error: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("Starting Konduktor server on %s:%d\n", cfg.Server.Host, cfg.Server.Port)
|
||||
|
||||
if err := srv.Run(); err != nil {
|
||||
return fmt.Errorf("server startup error: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var (
|
||||
Version = "0.1.0"
|
||||
BuildTime = "unknown"
|
||||
GitCommit = "unknown"
|
||||
)
|
||||
|
||||
func main() {
|
||||
rootCmd := &cobra.Command{
|
||||
Use: "konduktorctl",
|
||||
Short: "Konduktorctl - Service management CLI",
|
||||
Long: `Konduktorctl is a CLI tool for managing Konduktor services.`,
|
||||
Version: fmt.Sprintf("%s (commit: %s, built: %s)", Version, GitCommit, BuildTime),
|
||||
}
|
||||
|
||||
rootCmd.AddCommand(
|
||||
newUpCmd(),
|
||||
newDownCmd(),
|
||||
newStatusCmd(),
|
||||
newLogsCmd(),
|
||||
newHealthCmd(),
|
||||
newScaleCmd(),
|
||||
newConfigCmd(),
|
||||
newInitCmd(),
|
||||
newTopCmd(),
|
||||
)
|
||||
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func newUpCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "up [service...]",
|
||||
Short: "Start services",
|
||||
Long: `Start one or more services. If no service is specified, all services are started.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Starting services...")
|
||||
// TODO: Implement service start logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().BoolP("detach", "d", false, "Run in background")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newDownCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "down [service...]",
|
||||
Short: "Stop services",
|
||||
Long: `Stop one or more services. If no service is specified, all services are stopped.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Stopping services...")
|
||||
// TODO: Implement service stop logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newStatusCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "status [service...]",
|
||||
Short: "Show service status",
|
||||
Long: `Show the status of one or more services.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Service status:")
|
||||
// TODO: Implement status display logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newLogsCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "logs [service]",
|
||||
Short: "View service logs",
|
||||
Long: `View logs for a specific service.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Fetching logs...")
|
||||
// TODO: Implement logs viewing logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().BoolP("follow", "f", false, "Follow log output")
|
||||
cmd.Flags().IntP("tail", "n", 100, "Number of lines to show from the end")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newHealthCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "health",
|
||||
Short: "Check service health",
|
||||
Long: `Check the health status of all services.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Health check:")
|
||||
// TODO: Implement health check logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newScaleCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "scale <service>=<count>",
|
||||
Short: "Scale a service",
|
||||
Long: `Scale a service to a specific number of instances.`,
|
||||
Args: cobra.MinimumNArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Printf("Scaling: %v\n", args)
|
||||
// TODO: Implement scaling logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newConfigCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "config",
|
||||
Short: "Manage configuration",
|
||||
Long: `View and validate configuration.`,
|
||||
}
|
||||
|
||||
cmd.AddCommand(&cobra.Command{
|
||||
Use: "show",
|
||||
Short: "Show current configuration",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Current configuration:")
|
||||
// TODO: Implement config show logic
|
||||
return nil
|
||||
},
|
||||
})
|
||||
|
||||
cmd.AddCommand(&cobra.Command{
|
||||
Use: "validate",
|
||||
Short: "Validate configuration file",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Validating configuration...")
|
||||
// TODO: Implement config validation logic
|
||||
return nil
|
||||
},
|
||||
})
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newInitCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "init",
|
||||
Short: "Initialize a new project",
|
||||
Long: `Create a new Konduktor project with default configuration.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Initializing new project...")
|
||||
// TODO: Implement init logic
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newTopCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "top",
|
||||
Short: "Display running processes",
|
||||
Long: `Display real-time view of running processes and resource usage.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
fmt.Println("Process monitor:")
|
||||
// TODO: Implement top-like display
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
module github.com/konduktor/konduktor
|
||||
|
||||
go 1.23.0
|
||||
|
||||
toolchain go1.24.2
|
||||
|
||||
require (
|
||||
github.com/spf13/cobra v1.10.2
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/spf13/pflag v1.0.10 // indirect
|
||||
go.uber.org/multierr v1.10.0 // indirect
|
||||
go.uber.org/zap v1.27.1 // indirect
|
||||
gopkg.in/natefinch/lumberjack.v2 v2.2.1 // indirect
|
||||
)
|
||||
@@ -0,0 +1,20 @@
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
|
||||
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
|
||||
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
|
||||
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk=
|
||||
github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ=
|
||||
go.uber.org/multierr v1.10.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
|
||||
go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc=
|
||||
gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
@@ -0,0 +1,134 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
HTTP HTTPConfig `yaml:"http"`
|
||||
Server ServerConfig `yaml:"server"`
|
||||
SSL SSLConfig `yaml:"ssl"`
|
||||
Logging LoggingConfig `yaml:"logging"`
|
||||
Extensions []ExtensionConfig `yaml:"extensions"`
|
||||
}
|
||||
|
||||
type HTTPConfig struct {
|
||||
StaticDir string `yaml:"static_dir"`
|
||||
TemplatesDir string `yaml:"templates_dir"`
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
Host string `yaml:"host"`
|
||||
Port int `yaml:"port"`
|
||||
Backlog int `yaml:"backlog"`
|
||||
DefaultRoot bool `yaml:"default_root"`
|
||||
ProxyTimeout time.Duration `yaml:"proxy_timeout"`
|
||||
RedirectInstructions map[string]string `yaml:"redirect_instructions"`
|
||||
}
|
||||
|
||||
type SSLConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
CertFile string `yaml:"cert_file"`
|
||||
KeyFile string `yaml:"key_file"`
|
||||
}
|
||||
|
||||
type LoggingConfig struct {
|
||||
Level string `yaml:"level"`
|
||||
ConsoleOutput bool `yaml:"console_output"`
|
||||
Format LogFormatConfig `yaml:"format"`
|
||||
Console *ConsoleLogConfig `yaml:"console"`
|
||||
Files []FileLogConfig `yaml:"files"`
|
||||
}
|
||||
|
||||
type LogFormatConfig struct {
|
||||
Type string `yaml:"type"`
|
||||
UseColors bool `yaml:"use_colors"`
|
||||
ShowModule bool `yaml:"show_module"`
|
||||
TimestampFormat string `yaml:"timestamp_format"`
|
||||
}
|
||||
|
||||
type ConsoleLogConfig struct {
|
||||
Format LogFormatConfig `yaml:"format"`
|
||||
Level string `yaml:"level"`
|
||||
}
|
||||
|
||||
type FileLogConfig struct {
|
||||
Path string `yaml:"path"`
|
||||
Level string `yaml:"level"`
|
||||
Loggers []string `yaml:"loggers"`
|
||||
Format LogFormatConfig `yaml:"format"`
|
||||
MaxBytes int64 `yaml:"max_bytes"`
|
||||
BackupCount int `yaml:"backup_count"`
|
||||
}
|
||||
|
||||
type ExtensionConfig struct {
|
||||
Type string `yaml:"type"`
|
||||
Config map[string]interface{} `yaml:"config"`
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cfg := Default()
|
||||
if err := yaml.Unmarshal(data, cfg); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse config: %w", err)
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func Default() *Config {
|
||||
return &Config{
|
||||
HTTP: HTTPConfig{
|
||||
StaticDir: "./static",
|
||||
TemplatesDir: "./templates",
|
||||
},
|
||||
Server: ServerConfig{
|
||||
Host: "0.0.0.0",
|
||||
Port: 8080,
|
||||
Backlog: 5,
|
||||
DefaultRoot: false,
|
||||
ProxyTimeout: 30 * time.Second,
|
||||
},
|
||||
SSL: SSLConfig{
|
||||
Enabled: false,
|
||||
CertFile: "./ssl/cert.pem",
|
||||
KeyFile: "./ssl/key.pem",
|
||||
},
|
||||
Logging: LoggingConfig{
|
||||
Level: "INFO",
|
||||
ConsoleOutput: true,
|
||||
Format: LogFormatConfig{
|
||||
Type: "standard",
|
||||
UseColors: true,
|
||||
ShowModule: true,
|
||||
TimestampFormat: "2006-01-02 15:04:05",
|
||||
},
|
||||
},
|
||||
Extensions: []ExtensionConfig{},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
if c.Server.Port < 1 || c.Server.Port > 65535 {
|
||||
return fmt.Errorf("invalid port: %d", c.Server.Port)
|
||||
}
|
||||
|
||||
if c.SSL.Enabled {
|
||||
if c.SSL.CertFile == "" {
|
||||
return fmt.Errorf("SSL enabled but cert_file not specified")
|
||||
}
|
||||
if c.SSL.KeyFile == "" {
|
||||
return fmt.Errorf("SSL enabled but key_file not specified")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDefault(t *testing.T) {
|
||||
cfg := Default()
|
||||
|
||||
if cfg.Server.Host != "0.0.0.0" {
|
||||
t.Errorf("Expected host 0.0.0.0, got %s", cfg.Server.Host)
|
||||
}
|
||||
|
||||
if cfg.Server.Port != 8080 {
|
||||
t.Errorf("Expected port 8080, got %d", cfg.Server.Port)
|
||||
}
|
||||
|
||||
if cfg.SSL.Enabled {
|
||||
t.Error("Expected SSL to be disabled by default")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
modify func(*Config)
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid default config",
|
||||
modify: func(c *Config) {},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid port - too low",
|
||||
modify: func(c *Config) {
|
||||
c.Server.Port = 0
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid port - too high",
|
||||
modify: func(c *Config) {
|
||||
c.Server.Port = 70000
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "SSL enabled without cert",
|
||||
modify: func(c *Config) {
|
||||
c.SSL.Enabled = true
|
||||
c.SSL.CertFile = ""
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "SSL enabled without key",
|
||||
modify: func(c *Config) {
|
||||
c.SSL.Enabled = true
|
||||
c.SSL.CertFile = "cert.pem"
|
||||
c.SSL.KeyFile = ""
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cfg := Default()
|
||||
tt.modify(cfg)
|
||||
|
||||
err := cfg.Validate()
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoad(t *testing.T) {
|
||||
// Create temporary config file
|
||||
content := `
|
||||
server:
|
||||
host: 127.0.0.1
|
||||
port: 3000
|
||||
|
||||
logging:
|
||||
level: DEBUG
|
||||
`
|
||||
tmpfile, err := os.CreateTemp("", "config-*.yaml")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.Remove(tmpfile.Name())
|
||||
|
||||
if _, err := tmpfile.Write([]byte(content)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := tmpfile.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cfg, err := Load(tmpfile.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Server.Host != "127.0.0.1" {
|
||||
t.Errorf("Expected host 127.0.0.1, got %s", cfg.Server.Host)
|
||||
}
|
||||
|
||||
if cfg.Server.Port != 3000 {
|
||||
t.Errorf("Expected port 3000, got %d", cfg.Server.Port)
|
||||
}
|
||||
|
||||
if cfg.Logging.Level != "DEBUG" {
|
||||
t.Errorf("Expected level DEBUG, got %s", cfg.Logging.Level)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadNotFound(t *testing.T) {
|
||||
_, err := Load("/nonexistent/config.yaml")
|
||||
if err == nil {
|
||||
t.Error("Expected error for non-existent file")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,466 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
type CachingExtension struct {
|
||||
BaseExtension
|
||||
cache map[string]*cacheEntry
|
||||
cachePatterns []*cachePattern
|
||||
defaultTTL time.Duration
|
||||
maxSize int
|
||||
currentSize int
|
||||
mu sync.RWMutex
|
||||
|
||||
hits int64
|
||||
misses int64
|
||||
}
|
||||
|
||||
type cacheEntry struct {
|
||||
key string
|
||||
body []byte
|
||||
headers http.Header
|
||||
statusCode int
|
||||
contentType string
|
||||
createdAt time.Time
|
||||
expiresAt time.Time
|
||||
size int
|
||||
}
|
||||
|
||||
type cachePattern struct {
|
||||
pattern *regexp.Regexp
|
||||
ttl time.Duration
|
||||
methods []string
|
||||
}
|
||||
|
||||
type CachingConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
DefaultTTL string `yaml:"default_ttl"` // e.g., "5m", "1h"
|
||||
MaxSizeMB int `yaml:"max_size_mb"` // Max cache size in MB
|
||||
CachePatterns []PatternConfig `yaml:"cache_patterns"` // Patterns to cache
|
||||
}
|
||||
|
||||
type PatternConfig struct {
|
||||
Pattern string `yaml:"pattern"` // Regex pattern
|
||||
TTL string `yaml:"ttl"` // TTL for this pattern
|
||||
Methods []string `yaml:"methods"` // HTTP methods to cache (default: GET)
|
||||
}
|
||||
|
||||
func NewCachingExtension(config map[string]interface{}, logger *logging.Logger) (Extension, error) {
|
||||
ext := &CachingExtension{
|
||||
BaseExtension: NewBaseExtension("caching", 20, logger),
|
||||
cache: make(map[string]*cacheEntry),
|
||||
cachePatterns: make([]*cachePattern, 0),
|
||||
defaultTTL: 5 * time.Minute,
|
||||
maxSize: 100 * 1024 * 1024, // 100MB default
|
||||
}
|
||||
|
||||
if ttl, ok := config["default_ttl"].(string); ok {
|
||||
if duration, err := time.ParseDuration(ttl); err == nil {
|
||||
ext.defaultTTL = duration
|
||||
}
|
||||
}
|
||||
|
||||
if maxSize, ok := config["max_size_mb"].(int); ok {
|
||||
ext.maxSize = maxSize * 1024 * 1024
|
||||
} else if maxSizeFloat, ok := config["max_size_mb"].(float64); ok {
|
||||
ext.maxSize = int(maxSizeFloat) * 1024 * 1024
|
||||
}
|
||||
|
||||
if patterns, ok := config["cache_patterns"].([]interface{}); ok {
|
||||
for _, p := range patterns {
|
||||
if patternCfg, ok := p.(map[string]interface{}); ok {
|
||||
pattern := &cachePattern{
|
||||
ttl: ext.defaultTTL,
|
||||
methods: []string{"GET"},
|
||||
}
|
||||
|
||||
if patternStr, ok := patternCfg["pattern"].(string); ok {
|
||||
re, err := regexp.Compile(patternStr)
|
||||
if err != nil {
|
||||
logger.Error("Invalid cache pattern", "pattern", patternStr, "error", err)
|
||||
continue
|
||||
}
|
||||
pattern.pattern = re
|
||||
}
|
||||
|
||||
if ttl, ok := patternCfg["ttl"].(string); ok {
|
||||
if duration, err := time.ParseDuration(ttl); err == nil {
|
||||
pattern.ttl = duration
|
||||
}
|
||||
}
|
||||
|
||||
if methods, ok := patternCfg["methods"].([]interface{}); ok {
|
||||
pattern.methods = make([]string, 0)
|
||||
for _, m := range methods {
|
||||
if method, ok := m.(string); ok {
|
||||
pattern.methods = append(pattern.methods, strings.ToUpper(method))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ext.cachePatterns = append(ext.cachePatterns, pattern)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
go ext.cleanupLoop()
|
||||
|
||||
return ext, nil
|
||||
}
|
||||
|
||||
func (e *CachingExtension) ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (bool, error) {
|
||||
if !e.shouldCache(r) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
key := e.cacheKey(r)
|
||||
|
||||
e.mu.RLock()
|
||||
entry, exists := e.cache[key]
|
||||
e.mu.RUnlock()
|
||||
|
||||
if exists && time.Now().Before(entry.expiresAt) {
|
||||
e.mu.Lock()
|
||||
e.hits++
|
||||
e.mu.Unlock()
|
||||
|
||||
// Mark as cache hit to prevent setting X-Cache: MISS
|
||||
// Try to find cachingResponseWriter in the wrapper chain
|
||||
setCacheHitFlag(w)
|
||||
|
||||
for k, values := range entry.headers {
|
||||
for _, v := range values {
|
||||
w.Header().Add(k, v)
|
||||
}
|
||||
}
|
||||
w.Header().Set("X-Cache", "HIT")
|
||||
w.Header().Set("Content-Type", entry.contentType)
|
||||
w.WriteHeader(entry.statusCode)
|
||||
w.Write(entry.body)
|
||||
|
||||
e.logger.Debug("Cache hit", "key", key, "path", r.URL.Path)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
e.mu.Lock()
|
||||
e.misses++
|
||||
e.mu.Unlock()
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// setCacheHitFlag tries to find cachingResponseWriter and set cache hit flag
|
||||
func setCacheHitFlag(w http.ResponseWriter) {
|
||||
// Direct match
|
||||
if cw, ok := w.(*cachingResponseWriter); ok {
|
||||
cw.SetCacheHit()
|
||||
return
|
||||
}
|
||||
|
||||
// Try unwrapping
|
||||
type unwrapper interface {
|
||||
Unwrap() http.ResponseWriter
|
||||
}
|
||||
|
||||
for {
|
||||
if u, ok := w.(unwrapper); ok {
|
||||
w = u.Unwrap()
|
||||
if cw, ok := w.(*cachingResponseWriter); ok {
|
||||
cw.SetCacheHit()
|
||||
return
|
||||
}
|
||||
} else {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessResponse caches the response if applicable
|
||||
func (e *CachingExtension) ProcessResponse(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
||||
// Response caching is handled by the CachingResponseWriter
|
||||
// X-Cache header is set in the cachingResponseWriter.WriteHeader
|
||||
}
|
||||
|
||||
// WrapResponseWriter wraps the response writer to capture the response for caching
|
||||
func (e *CachingExtension) WrapResponseWriter(w http.ResponseWriter, r *http.Request) http.ResponseWriter {
|
||||
if !e.shouldCache(r) {
|
||||
return w
|
||||
}
|
||||
|
||||
return &cachingResponseWriter{
|
||||
ResponseWriter: w,
|
||||
ext: e,
|
||||
request: r,
|
||||
buffer: &bytes.Buffer{},
|
||||
}
|
||||
}
|
||||
|
||||
type cachingResponseWriter struct {
|
||||
http.ResponseWriter
|
||||
ext *CachingExtension
|
||||
request *http.Request
|
||||
buffer *bytes.Buffer
|
||||
statusCode int
|
||||
wroteHeader bool
|
||||
cacheHit bool // Flag to indicate if this was a cache hit
|
||||
}
|
||||
|
||||
func (cw *cachingResponseWriter) WriteHeader(code int) {
|
||||
if !cw.wroteHeader {
|
||||
cw.statusCode = code
|
||||
cw.wroteHeader = true
|
||||
// Set X-Cache: MISS header before writing headers (only if not a cache hit)
|
||||
if !cw.cacheHit {
|
||||
cw.ResponseWriter.Header().Set("X-Cache", "MISS")
|
||||
}
|
||||
cw.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
}
|
||||
|
||||
// SetCacheHit marks this response as a cache hit (to avoid setting X-Cache: MISS)
|
||||
func (cw *cachingResponseWriter) SetCacheHit() {
|
||||
cw.cacheHit = true
|
||||
}
|
||||
|
||||
func (cw *cachingResponseWriter) Write(b []byte) (int, error) {
|
||||
if !cw.wroteHeader {
|
||||
cw.WriteHeader(http.StatusOK)
|
||||
}
|
||||
|
||||
cw.buffer.Write(b)
|
||||
|
||||
return cw.ResponseWriter.Write(b)
|
||||
}
|
||||
|
||||
func (cw *cachingResponseWriter) Finalize() {
|
||||
if cw.statusCode < 200 || cw.statusCode >= 400 {
|
||||
return
|
||||
}
|
||||
|
||||
body := cw.buffer.Bytes()
|
||||
if len(body) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
key := cw.ext.cacheKey(cw.request)
|
||||
ttl := cw.ext.getTTL(cw.request)
|
||||
|
||||
entry := &cacheEntry{
|
||||
key: key,
|
||||
body: body,
|
||||
headers: cw.Header().Clone(),
|
||||
statusCode: cw.statusCode,
|
||||
contentType: cw.Header().Get("Content-Type"),
|
||||
createdAt: time.Now(),
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
size: len(body),
|
||||
}
|
||||
|
||||
cw.ext.store(entry)
|
||||
}
|
||||
|
||||
func (e *CachingExtension) shouldCache(r *http.Request) bool {
|
||||
path := r.URL.Path
|
||||
method := r.Method
|
||||
|
||||
for _, pattern := range e.cachePatterns {
|
||||
if pattern.pattern.MatchString(path) {
|
||||
for _, m := range pattern.methods {
|
||||
if m == method {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Default: only cache GET requests
|
||||
return method == "GET" && len(e.cachePatterns) == 0
|
||||
}
|
||||
|
||||
func (e *CachingExtension) cacheKey(r *http.Request) string {
|
||||
// Create cache key from method + URL + relevant headers
|
||||
h := sha256.New()
|
||||
h.Write([]byte(r.Method))
|
||||
h.Write([]byte(r.URL.String()))
|
||||
|
||||
// Include Accept-Encoding for vary
|
||||
if ae := r.Header.Get("Accept-Encoding"); ae != "" {
|
||||
h.Write([]byte(ae))
|
||||
}
|
||||
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
func (e *CachingExtension) getTTL(r *http.Request) time.Duration {
|
||||
path := r.URL.Path
|
||||
|
||||
for _, pattern := range e.cachePatterns {
|
||||
if pattern.pattern.MatchString(path) {
|
||||
return pattern.ttl
|
||||
}
|
||||
}
|
||||
|
||||
return e.defaultTTL
|
||||
}
|
||||
|
||||
func (e *CachingExtension) store(entry *cacheEntry) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
// Evict old entries if needed
|
||||
for e.currentSize+entry.size > e.maxSize && len(e.cache) > 0 {
|
||||
e.evictOldest()
|
||||
}
|
||||
|
||||
// Store new entry
|
||||
if existing, ok := e.cache[entry.key]; ok {
|
||||
e.currentSize -= existing.size
|
||||
}
|
||||
|
||||
e.cache[entry.key] = entry
|
||||
e.currentSize += entry.size
|
||||
|
||||
e.logger.Debug("Cached response",
|
||||
"key", entry.key[:16],
|
||||
"size", entry.size,
|
||||
"ttl", entry.expiresAt.Sub(entry.createdAt).String())
|
||||
}
|
||||
|
||||
func (e *CachingExtension) evictOldest() {
|
||||
var oldestKey string
|
||||
var oldestTime time.Time
|
||||
|
||||
for key, entry := range e.cache {
|
||||
if oldestKey == "" || entry.createdAt.Before(oldestTime) {
|
||||
oldestKey = key
|
||||
oldestTime = entry.createdAt
|
||||
}
|
||||
}
|
||||
|
||||
if oldestKey != "" {
|
||||
entry := e.cache[oldestKey]
|
||||
e.currentSize -= entry.size
|
||||
delete(e.cache, oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CachingExtension) cleanupLoop() {
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for range ticker.C {
|
||||
e.cleanupExpired()
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CachingExtension) cleanupExpired() {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, entry := range e.cache {
|
||||
if now.After(entry.expiresAt) {
|
||||
e.currentSize -= entry.size
|
||||
delete(e.cache, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CachingExtension) Invalidate(key string) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
if entry, ok := e.cache[key]; ok {
|
||||
e.currentSize -= entry.size
|
||||
delete(e.cache, key)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidatePattern removes all entries matching a pattern, unlike Invalidate
|
||||
func (e *CachingExtension) InvalidatePattern(pattern string) error {
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
for key, entry := range e.cache {
|
||||
if re.MatchString(key) {
|
||||
e.currentSize -= entry.size
|
||||
delete(e.cache, key)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Clear removes all entries from cache
|
||||
func (e *CachingExtension) Clear() {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
e.cache = make(map[string]*cacheEntry)
|
||||
e.currentSize = 0
|
||||
}
|
||||
|
||||
// GetMetrics returns caching metrics
|
||||
func (e *CachingExtension) GetMetrics() map[string]interface{} {
|
||||
e.mu.RLock()
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
hitRate := float64(0)
|
||||
total := e.hits + e.misses
|
||||
if total > 0 {
|
||||
hitRate = float64(e.hits) / float64(total) * 100
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"entries": len(e.cache),
|
||||
"size_bytes": e.currentSize,
|
||||
"max_size": e.maxSize,
|
||||
"hits": e.hits,
|
||||
"misses": e.misses,
|
||||
"hit_rate": hitRate,
|
||||
"patterns": len(e.cachePatterns),
|
||||
"default_ttl": e.defaultTTL.String(),
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup stops the cleanup goroutine
|
||||
func (e *CachingExtension) Cleanup() error {
|
||||
e.Clear()
|
||||
return nil
|
||||
}
|
||||
|
||||
// CacheReader wraps an io.ReadCloser to cache the body
|
||||
type CacheReader struct {
|
||||
io.ReadCloser
|
||||
buffer *bytes.Buffer
|
||||
}
|
||||
|
||||
func (cr *CacheReader) Read(p []byte) (int, error) {
|
||||
n, err := cr.ReadCloser.Read(p)
|
||||
if n > 0 {
|
||||
cr.buffer.Write(p[:n])
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (cr *CacheReader) GetBody() []byte {
|
||||
return cr.buffer.Bytes()
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
// Extension is the interface that all extensions must implement
|
||||
type Extension interface {
|
||||
// Name returns the unique name of the extension
|
||||
Name() string
|
||||
|
||||
// Initialize is called when the extension is loaded
|
||||
Initialize() error
|
||||
|
||||
// ProcessRequest processes an incoming request before routing.
|
||||
// Returns:
|
||||
// - response: if non-nil, the request is handled and no further processing occurs
|
||||
// - handled: if true, the request was handled by this extension
|
||||
// - err: any error that occurred
|
||||
ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (handled bool, err error)
|
||||
|
||||
// ProcessResponse is called after the response is generated but before it's sent.
|
||||
// Extensions can modify the response here.
|
||||
ProcessResponse(ctx context.Context, w http.ResponseWriter, r *http.Request)
|
||||
|
||||
// Cleanup is called when the extension is being unloaded
|
||||
Cleanup() error
|
||||
|
||||
// Enabled returns whether the extension is currently enabled
|
||||
Enabled() bool
|
||||
|
||||
// SetEnabled enables or disables the extension
|
||||
SetEnabled(enabled bool)
|
||||
|
||||
// Priority returns the extension's priority (lower = earlier execution)
|
||||
Priority() int
|
||||
}
|
||||
|
||||
// BaseExtension provides a default implementation for common Extension methods
|
||||
type BaseExtension struct {
|
||||
name string
|
||||
enabled bool
|
||||
priority int
|
||||
logger *logging.Logger
|
||||
}
|
||||
|
||||
// NewBaseExtension creates a new BaseExtension
|
||||
func NewBaseExtension(name string, priority int, logger *logging.Logger) BaseExtension {
|
||||
return BaseExtension{
|
||||
name: name,
|
||||
enabled: true,
|
||||
priority: priority,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// Name returns the extension name
|
||||
func (b *BaseExtension) Name() string {
|
||||
return b.name
|
||||
}
|
||||
|
||||
// Enabled returns whether the extension is enabled
|
||||
func (b *BaseExtension) Enabled() bool {
|
||||
return b.enabled
|
||||
}
|
||||
|
||||
// SetEnabled sets the enabled state
|
||||
func (b *BaseExtension) SetEnabled(enabled bool) {
|
||||
b.enabled = enabled
|
||||
}
|
||||
|
||||
// Priority returns the extension priority
|
||||
func (b *BaseExtension) Priority() int {
|
||||
return b.priority
|
||||
}
|
||||
|
||||
// Initialize default implementation (no-op)
|
||||
func (b *BaseExtension) Initialize() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cleanup default implementation (no-op)
|
||||
func (b *BaseExtension) Cleanup() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProcessRequest default implementation (pass-through)
|
||||
func (b *BaseExtension) ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// ProcessResponse default implementation (no-op)
|
||||
func (b *BaseExtension) ProcessResponse(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
// Logger returns the extension's logger
|
||||
func (b *BaseExtension) Logger() *logging.Logger {
|
||||
return b.logger
|
||||
}
|
||||
|
||||
// ExtensionConfig holds configuration for creating extensions
|
||||
type ExtensionConfig struct {
|
||||
Type string
|
||||
Config map[string]interface{}
|
||||
}
|
||||
|
||||
// ExtensionFactory is a function that creates an extension from config
|
||||
type ExtensionFactory func(config map[string]interface{}, logger *logging.Logger) (Extension, error)
|
||||
|
||||
// ResponseWriterWrapper is an optional interface that extensions can implement
|
||||
// to wrap the response writer for capturing/modifying responses
|
||||
type ResponseWriterWrapper interface {
|
||||
WrapResponseWriter(w http.ResponseWriter, r *http.Request) http.ResponseWriter
|
||||
}
|
||||
|
||||
// ResponseFinalizer is an optional interface for response writers that need
|
||||
// to perform finalization after the response is written (e.g., caching)
|
||||
type ResponseFinalizer interface {
|
||||
Finalize()
|
||||
}
|
||||
@@ -0,0 +1,271 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sort"
|
||||
"sync"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
// Manager manages all loaded extensions
|
||||
type Manager struct {
|
||||
extensions []Extension
|
||||
registry map[string]ExtensionFactory
|
||||
logger *logging.Logger
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewManager creates a new extension manager
|
||||
func NewManager(logger *logging.Logger) *Manager {
|
||||
m := &Manager{
|
||||
extensions: make([]Extension, 0),
|
||||
registry: make(map[string]ExtensionFactory),
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
// Register built-in extensions
|
||||
m.RegisterFactory("routing", NewRoutingExtension)
|
||||
m.RegisterFactory("security", NewSecurityExtension)
|
||||
m.RegisterFactory("caching", NewCachingExtension)
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// RegisterFactory registers an extension factory
|
||||
func (m *Manager) RegisterFactory(name string, factory ExtensionFactory) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.registry[name] = factory
|
||||
}
|
||||
|
||||
// LoadExtension loads an extension by type and config
|
||||
func (m *Manager) LoadExtension(extType string, config map[string]interface{}) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
factory, ok := m.registry[extType]
|
||||
if !ok {
|
||||
return fmt.Errorf("unknown extension type: %s", extType)
|
||||
}
|
||||
|
||||
ext, err := factory(config, m.logger)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create extension %s: %w", extType, err)
|
||||
}
|
||||
|
||||
if err := ext.Initialize(); err != nil {
|
||||
return fmt.Errorf("failed to initialize extension %s: %w", extType, err)
|
||||
}
|
||||
|
||||
m.extensions = append(m.extensions, ext)
|
||||
|
||||
// Sort by priority (lower first)
|
||||
sort.Slice(m.extensions, func(i, j int) bool {
|
||||
return m.extensions[i].Priority() < m.extensions[j].Priority()
|
||||
})
|
||||
|
||||
m.logger.Info("Loaded extension", "type", extType, "name", ext.Name(), "priority", ext.Priority())
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddExtension adds a pre-created extension
|
||||
func (m *Manager) AddExtension(ext Extension) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if err := ext.Initialize(); err != nil {
|
||||
return fmt.Errorf("failed to initialize extension %s: %w", ext.Name(), err)
|
||||
}
|
||||
|
||||
m.extensions = append(m.extensions, ext)
|
||||
|
||||
// Sort by priority
|
||||
sort.Slice(m.extensions, func(i, j int) bool {
|
||||
return m.extensions[i].Priority() < m.extensions[j].Priority()
|
||||
})
|
||||
|
||||
m.logger.Info("Added extension", "name", ext.Name(), "priority", ext.Priority())
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProcessRequest runs all extensions' ProcessRequest in order
|
||||
// Returns true if any extension handled the request
|
||||
func (m *Manager) ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (bool, error) {
|
||||
m.mu.RLock()
|
||||
extensions := m.extensions
|
||||
m.mu.RUnlock()
|
||||
|
||||
for _, ext := range extensions {
|
||||
if !ext.Enabled() {
|
||||
continue
|
||||
}
|
||||
|
||||
handled, err := ext.ProcessRequest(ctx, w, r)
|
||||
if err != nil {
|
||||
m.logger.Error("Extension error", "extension", ext.Name(), "error", err)
|
||||
// Continue to next extension on error
|
||||
continue
|
||||
}
|
||||
|
||||
if handled {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// ProcessResponse runs all extensions' ProcessResponse in reverse order
|
||||
func (m *Manager) ProcessResponse(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
||||
m.mu.RLock()
|
||||
extensions := m.extensions
|
||||
m.mu.RUnlock()
|
||||
|
||||
// Process in reverse order for response
|
||||
for i := len(extensions) - 1; i >= 0; i-- {
|
||||
ext := extensions[i]
|
||||
if !ext.Enabled() {
|
||||
continue
|
||||
}
|
||||
|
||||
ext.ProcessResponse(ctx, w, r)
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup cleans up all extensions
|
||||
func (m *Manager) Cleanup() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for _, ext := range m.extensions {
|
||||
if err := ext.Cleanup(); err != nil {
|
||||
m.logger.Error("Extension cleanup error", "extension", ext.Name(), "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
m.extensions = nil
|
||||
}
|
||||
|
||||
// GetExtension returns an extension by name
|
||||
func (m *Manager) GetExtension(name string) Extension {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for _, ext := range m.extensions {
|
||||
if ext.Name() == name {
|
||||
return ext
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Extensions returns all loaded extensions
|
||||
func (m *Manager) Extensions() []Extension {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
result := make([]Extension, len(m.extensions))
|
||||
copy(result, m.extensions)
|
||||
return result
|
||||
}
|
||||
|
||||
// Handler returns an http.Handler that processes requests through all extensions
|
||||
func (m *Manager) Handler(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Wrap response writer through all extensions that support it
|
||||
// Process in reverse priority order so highest priority wrapper is outermost
|
||||
wrappedWriter := w
|
||||
var finalizers []ResponseFinalizer
|
||||
|
||||
m.mu.RLock()
|
||||
extensions := m.extensions
|
||||
m.mu.RUnlock()
|
||||
|
||||
// Wrap response writer (lowest priority first, so they wrap in correct order)
|
||||
for _, ext := range extensions {
|
||||
if !ext.Enabled() {
|
||||
continue
|
||||
}
|
||||
if wrapper, ok := ext.(ResponseWriterWrapper); ok {
|
||||
wrappedWriter = wrapper.WrapResponseWriter(wrappedWriter, r)
|
||||
// Check if the wrapped writer implements Finalizer
|
||||
if finalizer, ok := wrappedWriter.(ResponseFinalizer); ok {
|
||||
finalizers = append(finalizers, finalizer)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Create response wrapper to capture status code
|
||||
responseWrapper := newResponseWrapper(wrappedWriter)
|
||||
|
||||
// Process request through extensions
|
||||
handled, err := m.ProcessRequest(ctx, responseWrapper, r)
|
||||
if err != nil {
|
||||
m.logger.Error("Error processing request", "error", err)
|
||||
}
|
||||
|
||||
if handled {
|
||||
// Extension handled the request, process response
|
||||
m.ProcessResponse(ctx, responseWrapper, r)
|
||||
// Finalize all response writers
|
||||
for i := len(finalizers) - 1; i >= 0; i-- {
|
||||
finalizers[i].Finalize()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// No extension handled, pass to next handler
|
||||
next.ServeHTTP(responseWrapper, r)
|
||||
|
||||
// Process response
|
||||
m.ProcessResponse(ctx, responseWrapper, r)
|
||||
|
||||
// Finalize all response writers
|
||||
for i := len(finalizers) - 1; i >= 0; i-- {
|
||||
finalizers[i].Finalize()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// responseWrapper wraps http.ResponseWriter to allow response modification
|
||||
type responseWrapper struct {
|
||||
http.ResponseWriter
|
||||
statusCode int
|
||||
written bool
|
||||
}
|
||||
|
||||
func newResponseWrapper(w http.ResponseWriter) *responseWrapper {
|
||||
return &responseWrapper{
|
||||
ResponseWriter: w,
|
||||
statusCode: http.StatusOK,
|
||||
}
|
||||
}
|
||||
|
||||
func (rw *responseWrapper) WriteHeader(code int) {
|
||||
if !rw.written {
|
||||
rw.statusCode = code
|
||||
rw.ResponseWriter.WriteHeader(code)
|
||||
rw.written = true
|
||||
}
|
||||
}
|
||||
|
||||
func (rw *responseWrapper) Write(b []byte) (int, error) {
|
||||
if !rw.written {
|
||||
rw.WriteHeader(http.StatusOK)
|
||||
}
|
||||
return rw.ResponseWriter.Write(b)
|
||||
}
|
||||
|
||||
func (rw *responseWrapper) StatusCode() int {
|
||||
return rw.statusCode
|
||||
}
|
||||
|
||||
// Unwrap returns the underlying ResponseWriter (for type assertions)
|
||||
func (rw *responseWrapper) Unwrap() http.ResponseWriter {
|
||||
return rw.ResponseWriter
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
func newTestLogger() *logging.Logger {
|
||||
logger, _ := logging.New(logging.Config{Level: "DEBUG"})
|
||||
return logger
|
||||
}
|
||||
|
||||
func TestNewManager(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
if manager == nil {
|
||||
t.Fatal("Expected manager, got nil")
|
||||
}
|
||||
|
||||
// Check built-in factories are registered
|
||||
if _, ok := manager.registry["routing"]; !ok {
|
||||
t.Error("Expected routing factory to be registered")
|
||||
}
|
||||
if _, ok := manager.registry["security"]; !ok {
|
||||
t.Error("Expected security factory to be registered")
|
||||
}
|
||||
if _, ok := manager.registry["caching"]; !ok {
|
||||
t.Error("Expected caching factory to be registered")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_LoadExtension(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
err := manager.LoadExtension("security", map[string]interface{}{})
|
||||
if err != nil {
|
||||
t.Errorf("Failed to load security extension: %v", err)
|
||||
}
|
||||
|
||||
exts := manager.Extensions()
|
||||
if len(exts) != 1 {
|
||||
t.Errorf("Expected 1 extension, got %d", len(exts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_LoadExtension_Unknown(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
err := manager.LoadExtension("unknown", map[string]interface{}{})
|
||||
if err == nil {
|
||||
t.Error("Expected error for unknown extension type")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_GetExtension(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
manager.LoadExtension("security", map[string]interface{}{})
|
||||
|
||||
ext := manager.GetExtension("security")
|
||||
if ext == nil {
|
||||
t.Error("Expected to find security extension")
|
||||
}
|
||||
|
||||
ext = manager.GetExtension("nonexistent")
|
||||
if ext != nil {
|
||||
t.Error("Expected nil for nonexistent extension")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_ProcessRequest(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
// Load security extension with blocked IP
|
||||
manager.LoadExtension("security", map[string]interface{}{
|
||||
"blocked_ips": []interface{}{"192.168.1.1"},
|
||||
})
|
||||
|
||||
// Create test request
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.1:12345"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handled, err := manager.ProcessRequest(context.Background(), rr, req)
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if !handled {
|
||||
t.Error("Expected request to be handled (blocked)")
|
||||
}
|
||||
|
||||
if rr.Code != http.StatusForbidden {
|
||||
t.Errorf("Expected status 403, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_Handler(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
// Load routing extension with a simple route
|
||||
manager.LoadExtension("routing", map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"=/health": map[string]interface{}{
|
||||
"return": "200 OK",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
baseHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
})
|
||||
|
||||
handler := manager.Handler(baseHandler)
|
||||
|
||||
// Test health route
|
||||
req := httptest.NewRequest("GET", "/health", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_Priority(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
// Load extensions in any order
|
||||
manager.LoadExtension("routing", map[string]interface{}{}) // Priority 50
|
||||
manager.LoadExtension("security", map[string]interface{}{}) // Priority 10
|
||||
manager.LoadExtension("caching", map[string]interface{}{}) // Priority 20
|
||||
|
||||
exts := manager.Extensions()
|
||||
if len(exts) != 3 {
|
||||
t.Fatalf("Expected 3 extensions, got %d", len(exts))
|
||||
}
|
||||
|
||||
// Check order by priority
|
||||
if exts[0].Name() != "security" {
|
||||
t.Errorf("Expected security first, got %s", exts[0].Name())
|
||||
}
|
||||
if exts[1].Name() != "caching" {
|
||||
t.Errorf("Expected caching second, got %s", exts[1].Name())
|
||||
}
|
||||
if exts[2].Name() != "routing" {
|
||||
t.Errorf("Expected routing third, got %s", exts[2].Name())
|
||||
}
|
||||
}
|
||||
|
||||
func TestManager_Cleanup(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
manager := NewManager(logger)
|
||||
|
||||
manager.LoadExtension("security", map[string]interface{}{})
|
||||
manager.LoadExtension("routing", map[string]interface{}{})
|
||||
|
||||
manager.Cleanup()
|
||||
|
||||
exts := manager.Extensions()
|
||||
if len(exts) != 0 {
|
||||
t.Errorf("Expected 0 extensions after cleanup, got %d", len(exts))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,428 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
"github.com/konduktor/konduktor/internal/proxy"
|
||||
)
|
||||
|
||||
// RoutingExtension handles request routing based on patterns
|
||||
type RoutingExtension struct {
|
||||
BaseExtension
|
||||
exactRoutes map[string]RouteConfig
|
||||
regexRoutes []*regexRoute
|
||||
defaultRoute *RouteConfig
|
||||
staticDir string
|
||||
}
|
||||
|
||||
// RouteConfig holds configuration for a route
|
||||
type RouteConfig struct {
|
||||
ProxyPass string
|
||||
Root string
|
||||
IndexFile string
|
||||
Return string
|
||||
ContentType string
|
||||
Headers []string
|
||||
CacheControl string
|
||||
SPAFallback bool
|
||||
ExcludePatterns []string
|
||||
Timeout float64
|
||||
}
|
||||
|
||||
type regexRoute struct {
|
||||
pattern *regexp.Regexp
|
||||
config RouteConfig
|
||||
caseSensitive bool
|
||||
originalExpr string
|
||||
}
|
||||
|
||||
// NewRoutingExtension creates a new routing extension
|
||||
func NewRoutingExtension(config map[string]interface{}, logger *logging.Logger) (Extension, error) {
|
||||
ext := &RoutingExtension{
|
||||
BaseExtension: NewBaseExtension("routing", 50, logger), // Middle priority
|
||||
exactRoutes: make(map[string]RouteConfig),
|
||||
regexRoutes: make([]*regexRoute, 0),
|
||||
staticDir: "./static",
|
||||
}
|
||||
|
||||
logger.Debug("Routing extension config", "config", config)
|
||||
|
||||
// Parse regex_locations from config
|
||||
if locations, ok := config["regex_locations"].(map[string]interface{}); ok {
|
||||
logger.Debug("Found regex_locations", "count", len(locations))
|
||||
for pattern, routeCfg := range locations {
|
||||
logger.Debug("Adding route", "pattern", pattern)
|
||||
if rc, ok := routeCfg.(map[string]interface{}); ok {
|
||||
ext.addRoute(pattern, parseRouteConfig(rc))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
logger.Warn("No regex_locations found in config", "config_keys", getKeys(config))
|
||||
}
|
||||
|
||||
// Parse static_dir if provided
|
||||
if staticDir, ok := config["static_dir"].(string); ok {
|
||||
ext.staticDir = staticDir
|
||||
}
|
||||
|
||||
return ext, nil
|
||||
}
|
||||
|
||||
func parseRouteConfig(cfg map[string]interface{}) RouteConfig {
|
||||
rc := RouteConfig{
|
||||
IndexFile: "index.html",
|
||||
}
|
||||
|
||||
if v, ok := cfg["proxy_pass"].(string); ok {
|
||||
rc.ProxyPass = v
|
||||
}
|
||||
if v, ok := cfg["root"].(string); ok {
|
||||
rc.Root = v
|
||||
}
|
||||
if v, ok := cfg["index_file"].(string); ok {
|
||||
rc.IndexFile = v
|
||||
}
|
||||
if v, ok := cfg["return"].(string); ok {
|
||||
rc.Return = v
|
||||
}
|
||||
if v, ok := cfg["content_type"].(string); ok {
|
||||
rc.ContentType = v
|
||||
}
|
||||
if v, ok := cfg["cache_control"].(string); ok {
|
||||
rc.CacheControl = v
|
||||
}
|
||||
if v, ok := cfg["spa_fallback"].(bool); ok {
|
||||
rc.SPAFallback = v
|
||||
}
|
||||
if v, ok := cfg["timeout"].(float64); ok {
|
||||
rc.Timeout = v
|
||||
}
|
||||
|
||||
// Parse headers
|
||||
if headers, ok := cfg["headers"].([]interface{}); ok {
|
||||
for _, h := range headers {
|
||||
if header, ok := h.(string); ok {
|
||||
rc.Headers = append(rc.Headers, header)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Parse exclude_patterns
|
||||
if patterns, ok := cfg["exclude_patterns"].([]interface{}); ok {
|
||||
for _, p := range patterns {
|
||||
if pattern, ok := p.(string); ok {
|
||||
rc.ExcludePatterns = append(rc.ExcludePatterns, pattern)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return rc
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) addRoute(pattern string, config RouteConfig) {
|
||||
switch {
|
||||
case pattern == "__default__":
|
||||
e.defaultRoute = &config
|
||||
|
||||
case strings.HasPrefix(pattern, "="):
|
||||
// Exact match
|
||||
path := strings.TrimPrefix(pattern, "=")
|
||||
e.exactRoutes[path] = config
|
||||
|
||||
case strings.HasPrefix(pattern, "~*"):
|
||||
// Case-insensitive regex
|
||||
expr := strings.TrimPrefix(pattern, "~*")
|
||||
re, err := regexp.Compile("(?i)" + expr)
|
||||
if err != nil {
|
||||
e.logger.Error("Invalid regex pattern", "pattern", pattern, "error", err)
|
||||
return
|
||||
}
|
||||
e.regexRoutes = append(e.regexRoutes, ®exRoute{
|
||||
pattern: re,
|
||||
config: config,
|
||||
caseSensitive: false,
|
||||
originalExpr: expr,
|
||||
})
|
||||
|
||||
case strings.HasPrefix(pattern, "~"):
|
||||
// Case-sensitive regex
|
||||
expr := strings.TrimPrefix(pattern, "~")
|
||||
re, err := regexp.Compile(expr)
|
||||
if err != nil {
|
||||
e.logger.Error("Invalid regex pattern", "pattern", pattern, "error", err)
|
||||
return
|
||||
}
|
||||
e.regexRoutes = append(e.regexRoutes, ®exRoute{
|
||||
pattern: re,
|
||||
config: config,
|
||||
caseSensitive: true,
|
||||
originalExpr: expr,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessRequest handles the request routing
|
||||
func (e *RoutingExtension) ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (bool, error) {
|
||||
path := r.URL.Path
|
||||
|
||||
// 1. Check exact routes (ignore request path for proxy)
|
||||
if config, ok := e.exactRoutes[path]; ok {
|
||||
return e.handleRoute(w, r, config, nil, true)
|
||||
}
|
||||
|
||||
// 2. Check regex routes
|
||||
for _, route := range e.regexRoutes {
|
||||
match := route.pattern.FindStringSubmatch(path)
|
||||
if match != nil {
|
||||
params := make(map[string]string)
|
||||
names := route.pattern.SubexpNames()
|
||||
for i, name := range names {
|
||||
if i > 0 && name != "" && i < len(match) {
|
||||
params[name] = match[i]
|
||||
}
|
||||
}
|
||||
return e.handleRoute(w, r, route.config, params, false)
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Check default route
|
||||
if e.defaultRoute != nil {
|
||||
return e.handleRoute(w, r, *e.defaultRoute, nil, false)
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) handleRoute(w http.ResponseWriter, r *http.Request, config RouteConfig, params map[string]string, exactMatch bool) (bool, error) {
|
||||
// Handle "return" directive
|
||||
if config.Return != "" {
|
||||
return e.handleReturn(w, config)
|
||||
}
|
||||
|
||||
// Handle proxy_pass
|
||||
if config.ProxyPass != "" {
|
||||
return e.handleProxy(w, r, config, params, exactMatch)
|
||||
}
|
||||
|
||||
// Handle static files with root
|
||||
if config.Root != "" {
|
||||
return e.handleStatic(w, r, config)
|
||||
}
|
||||
|
||||
// Handle SPA fallback
|
||||
if config.SPAFallback {
|
||||
return e.handleSPAFallback(w, r, config)
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) handleReturn(w http.ResponseWriter, config RouteConfig) (bool, error) {
|
||||
parts := strings.SplitN(config.Return, " ", 2)
|
||||
statusCode := 200
|
||||
body := "OK"
|
||||
|
||||
if len(parts) >= 1 {
|
||||
switch parts[0] {
|
||||
case "200":
|
||||
statusCode = 200
|
||||
case "201":
|
||||
statusCode = 201
|
||||
case "301":
|
||||
statusCode = 301
|
||||
case "302":
|
||||
statusCode = 302
|
||||
case "400":
|
||||
statusCode = 400
|
||||
case "404":
|
||||
statusCode = 404
|
||||
case "500":
|
||||
statusCode = 500
|
||||
}
|
||||
}
|
||||
if len(parts) >= 2 {
|
||||
body = parts[1]
|
||||
}
|
||||
|
||||
contentType := "text/plain"
|
||||
if config.ContentType != "" {
|
||||
contentType = config.ContentType
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
w.WriteHeader(statusCode)
|
||||
w.Write([]byte(body))
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) handleProxy(w http.ResponseWriter, r *http.Request, config RouteConfig, params map[string]string, exactMatch bool) (bool, error) {
|
||||
target := config.ProxyPass
|
||||
|
||||
// Check if target URL contains parameter placeholders
|
||||
hasParams := strings.Contains(target, "{") && strings.Contains(target, "}")
|
||||
|
||||
// Substitute params in target URL
|
||||
for key, value := range params {
|
||||
target = strings.ReplaceAll(target, "{"+key+"}", value)
|
||||
}
|
||||
|
||||
// Create proxy config
|
||||
// IgnoreRequestPath=true when:
|
||||
// - exact match route (=/path)
|
||||
// - target URL had parameter substitutions (the target path is fully specified)
|
||||
proxyConfig := &proxy.Config{
|
||||
Target: target,
|
||||
Headers: make(map[string]string),
|
||||
IgnoreRequestPath: exactMatch || hasParams,
|
||||
}
|
||||
|
||||
// Set timeout if specified
|
||||
if config.Timeout > 0 {
|
||||
proxyConfig.Timeout = time.Duration(config.Timeout * float64(time.Second))
|
||||
}
|
||||
|
||||
// Parse headers
|
||||
clientIP := getClientIP(r)
|
||||
for _, header := range config.Headers {
|
||||
parts := strings.SplitN(header, ": ", 2)
|
||||
if len(parts) == 2 {
|
||||
value := parts[1]
|
||||
// Substitute params
|
||||
for key, pValue := range params {
|
||||
value = strings.ReplaceAll(value, "{"+key+"}", pValue)
|
||||
}
|
||||
// Substitute special variables
|
||||
value = strings.ReplaceAll(value, "$remote_addr", clientIP)
|
||||
proxyConfig.Headers[parts[0]] = value
|
||||
}
|
||||
}
|
||||
|
||||
p, err := proxy.New(proxyConfig, e.logger)
|
||||
if err != nil {
|
||||
e.logger.Error("Failed to create proxy", "target", target, "error", err)
|
||||
http.Error(w, "Bad Gateway", http.StatusBadGateway)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
p.ProxyRequest(w, r, params)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) handleStatic(w http.ResponseWriter, r *http.Request, config RouteConfig) (bool, error) {
|
||||
path := r.URL.Path
|
||||
|
||||
// Handle index file for root or directory paths
|
||||
if path == "/" || strings.HasSuffix(path, "/") {
|
||||
path = "/" + config.IndexFile
|
||||
}
|
||||
|
||||
// Get absolute path for root dir
|
||||
absRoot, err := filepath.Abs(config.Root)
|
||||
if err != nil {
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
filePath := filepath.Join(absRoot, filepath.Clean("/"+path))
|
||||
cleanPath := filepath.Clean(filePath)
|
||||
|
||||
// Prevent directory traversal
|
||||
if !strings.HasPrefix(cleanPath+string(filepath.Separator), absRoot+string(filepath.Separator)) {
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// Check if file exists
|
||||
if _, err := os.Stat(filePath); os.IsNotExist(err) {
|
||||
return false, nil // Let other handlers try
|
||||
}
|
||||
|
||||
// Set cache control header
|
||||
if config.CacheControl != "" {
|
||||
w.Header().Set("Cache-Control", config.CacheControl)
|
||||
}
|
||||
|
||||
// Set custom headers
|
||||
for _, header := range config.Headers {
|
||||
parts := strings.SplitN(header, ": ", 2)
|
||||
if len(parts) == 2 {
|
||||
w.Header().Set(parts[0], parts[1])
|
||||
}
|
||||
}
|
||||
|
||||
http.ServeFile(w, r, filePath)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (e *RoutingExtension) handleSPAFallback(w http.ResponseWriter, r *http.Request, config RouteConfig) (bool, error) {
|
||||
path := r.URL.Path
|
||||
|
||||
// Check exclude patterns
|
||||
for _, pattern := range config.ExcludePatterns {
|
||||
if strings.HasPrefix(path, pattern) {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
root := config.Root
|
||||
if root == "" {
|
||||
root = e.staticDir
|
||||
}
|
||||
|
||||
indexFile := config.IndexFile
|
||||
if indexFile == "" {
|
||||
indexFile = "index.html"
|
||||
}
|
||||
|
||||
filePath := filepath.Join(root, indexFile)
|
||||
if _, err := os.Stat(filePath); os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
http.ServeFile(w, r, filePath)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func getClientIP(r *http.Request) string {
|
||||
// Check X-Forwarded-For header first
|
||||
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
||||
parts := strings.Split(xff, ",")
|
||||
return strings.TrimSpace(parts[0])
|
||||
}
|
||||
|
||||
// Check X-Real-IP header
|
||||
if xri := r.Header.Get("X-Real-IP"); xri != "" {
|
||||
return xri
|
||||
}
|
||||
|
||||
// Fall back to RemoteAddr
|
||||
ip := r.RemoteAddr
|
||||
if idx := strings.LastIndex(ip, ":"); idx != -1 {
|
||||
ip = ip[:idx]
|
||||
}
|
||||
return ip
|
||||
}
|
||||
|
||||
// GetMetrics returns routing metrics
|
||||
func (e *RoutingExtension) GetMetrics() map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"exact_routes": len(e.exactRoutes),
|
||||
"regex_routes": len(e.regexRoutes),
|
||||
"has_default": e.defaultRoute != nil,
|
||||
}
|
||||
}
|
||||
|
||||
func getKeys(m map[string]interface{}) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
// SecurityExtension provides security features like IP filtering and security headers
|
||||
type SecurityExtension struct {
|
||||
BaseExtension
|
||||
allowedIPs map[string]bool
|
||||
blockedIPs map[string]bool
|
||||
allowedCIDRs []*net.IPNet
|
||||
blockedCIDRs []*net.IPNet
|
||||
securityHeaders map[string]string
|
||||
|
||||
// Rate limiting
|
||||
rateLimitEnabled bool
|
||||
rateLimitRequests int
|
||||
rateLimitWindow time.Duration
|
||||
rateLimitByIP map[string]*rateLimitEntry
|
||||
rateLimitMu sync.RWMutex
|
||||
}
|
||||
|
||||
type rateLimitEntry struct {
|
||||
count int
|
||||
resetTime time.Time
|
||||
}
|
||||
|
||||
// SecurityConfig holds security extension configuration
|
||||
type SecurityConfig struct {
|
||||
AllowedIPs []string `yaml:"allowed_ips"`
|
||||
BlockedIPs []string `yaml:"blocked_ips"`
|
||||
SecurityHeaders map[string]string `yaml:"security_headers"`
|
||||
RateLimit *RateLimitConfig `yaml:"rate_limit"`
|
||||
}
|
||||
|
||||
// RateLimitConfig holds rate limiting configuration
|
||||
type RateLimitConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
Requests int `yaml:"requests"`
|
||||
Window string `yaml:"window"` // e.g., "1m", "1h"
|
||||
}
|
||||
|
||||
// Default security headers
|
||||
var defaultSecurityHeaders = map[string]string{
|
||||
"X-Content-Type-Options": "nosniff",
|
||||
"X-Frame-Options": "DENY",
|
||||
"X-XSS-Protection": "1; mode=block",
|
||||
"Referrer-Policy": "strict-origin-when-cross-origin",
|
||||
}
|
||||
|
||||
// NewSecurityExtension creates a new security extension
|
||||
func NewSecurityExtension(config map[string]interface{}, logger *logging.Logger) (Extension, error) {
|
||||
ext := &SecurityExtension{
|
||||
BaseExtension: NewBaseExtension("security", 10, logger), // High priority (early execution)
|
||||
allowedIPs: make(map[string]bool),
|
||||
blockedIPs: make(map[string]bool),
|
||||
allowedCIDRs: make([]*net.IPNet, 0),
|
||||
blockedCIDRs: make([]*net.IPNet, 0),
|
||||
securityHeaders: make(map[string]string),
|
||||
rateLimitByIP: make(map[string]*rateLimitEntry),
|
||||
}
|
||||
|
||||
// Copy default security headers
|
||||
for k, v := range defaultSecurityHeaders {
|
||||
ext.securityHeaders[k] = v
|
||||
}
|
||||
|
||||
// Parse allowed_ips
|
||||
if allowedIPs, ok := config["allowed_ips"].([]interface{}); ok {
|
||||
for _, ip := range allowedIPs {
|
||||
if ipStr, ok := ip.(string); ok {
|
||||
ext.addAllowedIP(ipStr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Parse blocked_ips
|
||||
if blockedIPs, ok := config["blocked_ips"].([]interface{}); ok {
|
||||
for _, ip := range blockedIPs {
|
||||
if ipStr, ok := ip.(string); ok {
|
||||
ext.addBlockedIP(ipStr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Parse security_headers
|
||||
if headers, ok := config["security_headers"].(map[string]interface{}); ok {
|
||||
for k, v := range headers {
|
||||
if vStr, ok := v.(string); ok {
|
||||
ext.securityHeaders[k] = vStr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Parse rate_limit
|
||||
if rateLimit, ok := config["rate_limit"].(map[string]interface{}); ok {
|
||||
if enabled, ok := rateLimit["enabled"].(bool); ok && enabled {
|
||||
ext.rateLimitEnabled = true
|
||||
|
||||
if requests, ok := rateLimit["requests"].(int); ok {
|
||||
ext.rateLimitRequests = requests
|
||||
} else if requestsFloat, ok := rateLimit["requests"].(float64); ok {
|
||||
ext.rateLimitRequests = int(requestsFloat)
|
||||
} else {
|
||||
ext.rateLimitRequests = 100 // default
|
||||
}
|
||||
|
||||
if window, ok := rateLimit["window"].(string); ok {
|
||||
if duration, err := time.ParseDuration(window); err == nil {
|
||||
ext.rateLimitWindow = duration
|
||||
} else {
|
||||
ext.rateLimitWindow = time.Minute // default
|
||||
}
|
||||
} else {
|
||||
ext.rateLimitWindow = time.Minute
|
||||
}
|
||||
|
||||
logger.Info("Rate limiting enabled",
|
||||
"requests", ext.rateLimitRequests,
|
||||
"window", ext.rateLimitWindow.String())
|
||||
}
|
||||
}
|
||||
|
||||
return ext, nil
|
||||
}
|
||||
|
||||
func (e *SecurityExtension) addAllowedIP(ip string) {
|
||||
if strings.Contains(ip, "/") {
|
||||
// CIDR notation
|
||||
_, cidr, err := net.ParseCIDR(ip)
|
||||
if err == nil {
|
||||
e.allowedCIDRs = append(e.allowedCIDRs, cidr)
|
||||
}
|
||||
} else {
|
||||
e.allowedIPs[ip] = true
|
||||
}
|
||||
}
|
||||
|
||||
func (e *SecurityExtension) addBlockedIP(ip string) {
|
||||
if strings.Contains(ip, "/") {
|
||||
// CIDR notation
|
||||
_, cidr, err := net.ParseCIDR(ip)
|
||||
if err == nil {
|
||||
e.blockedCIDRs = append(e.blockedCIDRs, cidr)
|
||||
}
|
||||
} else {
|
||||
e.blockedIPs[ip] = true
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessRequest checks security rules
|
||||
func (e *SecurityExtension) ProcessRequest(ctx context.Context, w http.ResponseWriter, r *http.Request) (bool, error) {
|
||||
clientIP := getClientIP(r)
|
||||
parsedIP := net.ParseIP(clientIP)
|
||||
|
||||
// Check blocked IPs first
|
||||
if e.isBlocked(clientIP, parsedIP) {
|
||||
e.logger.Warn("Blocked request from IP", "ip", clientIP)
|
||||
http.Error(w, "403 Forbidden", http.StatusForbidden)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// Check allowed IPs (if configured, only these IPs are allowed)
|
||||
if len(e.allowedIPs) > 0 || len(e.allowedCIDRs) > 0 {
|
||||
if !e.isAllowed(clientIP, parsedIP) {
|
||||
e.logger.Warn("Access denied for IP", "ip", clientIP)
|
||||
http.Error(w, "403 Forbidden", http.StatusForbidden)
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Check rate limit
|
||||
if e.rateLimitEnabled {
|
||||
if !e.checkRateLimit(clientIP) {
|
||||
e.logger.Warn("Rate limit exceeded", "ip", clientIP)
|
||||
w.Header().Set("Retry-After", "60")
|
||||
http.Error(w, "429 Too Many Requests", http.StatusTooManyRequests)
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// ProcessResponse adds security headers to the response
|
||||
func (e *SecurityExtension) ProcessResponse(ctx context.Context, w http.ResponseWriter, r *http.Request) {
|
||||
for header, value := range e.securityHeaders {
|
||||
w.Header().Set(header, value)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *SecurityExtension) isBlocked(ip string, parsedIP net.IP) bool {
|
||||
// Check exact match
|
||||
if e.blockedIPs[ip] {
|
||||
return true
|
||||
}
|
||||
|
||||
// Check CIDR ranges
|
||||
if parsedIP != nil {
|
||||
for _, cidr := range e.blockedCIDRs {
|
||||
if cidr.Contains(parsedIP) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (e *SecurityExtension) isAllowed(ip string, parsedIP net.IP) bool {
|
||||
// Check exact match
|
||||
if e.allowedIPs[ip] {
|
||||
return true
|
||||
}
|
||||
|
||||
// Check CIDR ranges
|
||||
if parsedIP != nil {
|
||||
for _, cidr := range e.allowedCIDRs {
|
||||
if cidr.Contains(parsedIP) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (e *SecurityExtension) checkRateLimit(ip string) bool {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
entry, exists := e.rateLimitByIP[ip]
|
||||
|
||||
if !exists || now.After(entry.resetTime) {
|
||||
// Create new entry or reset expired one
|
||||
e.rateLimitByIP[ip] = &rateLimitEntry{
|
||||
count: 1,
|
||||
resetTime: now.Add(e.rateLimitWindow),
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// Increment counter
|
||||
entry.count++
|
||||
return entry.count <= e.rateLimitRequests
|
||||
}
|
||||
|
||||
// AddBlockedIP adds an IP to the blocked list at runtime
|
||||
func (e *SecurityExtension) AddBlockedIP(ip string) {
|
||||
e.addBlockedIP(ip)
|
||||
}
|
||||
|
||||
// RemoveBlockedIP removes an IP from the blocked list
|
||||
func (e *SecurityExtension) RemoveBlockedIP(ip string) {
|
||||
delete(e.blockedIPs, ip)
|
||||
}
|
||||
|
||||
// AddAllowedIP adds an IP to the allowed list at runtime
|
||||
func (e *SecurityExtension) AddAllowedIP(ip string) {
|
||||
e.addAllowedIP(ip)
|
||||
}
|
||||
|
||||
// RemoveAllowedIP removes an IP from the allowed list
|
||||
func (e *SecurityExtension) RemoveAllowedIP(ip string) {
|
||||
delete(e.allowedIPs, ip)
|
||||
}
|
||||
|
||||
// SetSecurityHeader sets or updates a security header
|
||||
func (e *SecurityExtension) SetSecurityHeader(name, value string) {
|
||||
e.securityHeaders[name] = value
|
||||
}
|
||||
|
||||
// GetMetrics returns security metrics
|
||||
func (e *SecurityExtension) GetMetrics() map[string]interface{} {
|
||||
e.rateLimitMu.RLock()
|
||||
activeRateLimits := len(e.rateLimitByIP)
|
||||
e.rateLimitMu.RUnlock()
|
||||
|
||||
return map[string]interface{}{
|
||||
"allowed_ips": len(e.allowedIPs),
|
||||
"allowed_cidrs": len(e.allowedCIDRs),
|
||||
"blocked_ips": len(e.blockedIPs),
|
||||
"blocked_cidrs": len(e.blockedCIDRs),
|
||||
"security_headers": len(e.securityHeaders),
|
||||
"rate_limit_enabled": e.rateLimitEnabled,
|
||||
"active_rate_limits": activeRateLimits,
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup cleans up rate limit entries periodically
|
||||
func (e *SecurityExtension) Cleanup() error {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for ip, entry := range e.rateLimitByIP {
|
||||
if now.After(entry.resetTime) {
|
||||
delete(e.rateLimitByIP, ip)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
package extension
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNewSecurityExtension(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, err := NewSecurityExtension(map[string]interface{}{}, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create security extension: %v", err)
|
||||
}
|
||||
|
||||
if ext.Name() != "security" {
|
||||
t.Errorf("Expected name 'security', got %s", ext.Name())
|
||||
}
|
||||
|
||||
if ext.Priority() != 10 {
|
||||
t.Errorf("Expected priority 10, got %d", ext.Priority())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_BlockedIP(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"blocked_ips": []interface{}{"192.168.1.100"},
|
||||
}, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.100:12345"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handled, err := ext.ProcessRequest(context.Background(), rr, req)
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if !handled {
|
||||
t.Error("Expected blocked request to be handled")
|
||||
}
|
||||
|
||||
if rr.Code != http.StatusForbidden {
|
||||
t.Errorf("Expected status 403, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_AllowedIP(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"allowed_ips": []interface{}{"192.168.1.50"},
|
||||
}, logger)
|
||||
|
||||
// Allowed IP
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.50:12345"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handled, _ := ext.ProcessRequest(context.Background(), rr, req)
|
||||
if handled {
|
||||
t.Error("Expected allowed IP to pass through")
|
||||
}
|
||||
|
||||
// Not allowed IP
|
||||
req = httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.51:12345"
|
||||
rr = httptest.NewRecorder()
|
||||
|
||||
handled, _ = ext.ProcessRequest(context.Background(), rr, req)
|
||||
if !handled {
|
||||
t.Error("Expected non-allowed IP to be blocked")
|
||||
}
|
||||
|
||||
if rr.Code != http.StatusForbidden {
|
||||
t.Errorf("Expected status 403, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_CIDR(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"blocked_ips": []interface{}{"10.0.0.0/8"},
|
||||
}, logger)
|
||||
|
||||
// IP in blocked CIDR
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "10.1.2.3:12345"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handled, _ := ext.ProcessRequest(context.Background(), rr, req)
|
||||
if !handled {
|
||||
t.Error("Expected IP in blocked CIDR to be blocked")
|
||||
}
|
||||
|
||||
// IP not in blocked CIDR
|
||||
req = httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.1:12345"
|
||||
rr = httptest.NewRecorder()
|
||||
|
||||
handled, _ = ext.ProcessRequest(context.Background(), rr, req)
|
||||
if handled {
|
||||
t.Error("Expected IP not in blocked CIDR to pass through")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_SecurityHeaders(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"security_headers": map[string]interface{}{
|
||||
"X-Custom-Header": "custom-value",
|
||||
},
|
||||
}, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
ext.ProcessResponse(context.Background(), rr, req)
|
||||
|
||||
// Check default headers
|
||||
if rr.Header().Get("X-Content-Type-Options") != "nosniff" {
|
||||
t.Error("Expected X-Content-Type-Options header")
|
||||
}
|
||||
|
||||
// Check custom header
|
||||
if rr.Header().Get("X-Custom-Header") != "custom-value" {
|
||||
t.Error("Expected custom header")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_RateLimit(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"rate_limit": map[string]interface{}{
|
||||
"enabled": true,
|
||||
"requests": 2,
|
||||
"window": "1m",
|
||||
},
|
||||
}, logger)
|
||||
|
||||
securityExt := ext.(*SecurityExtension)
|
||||
clientIP := "192.168.1.1"
|
||||
|
||||
// First request - should pass
|
||||
if !securityExt.checkRateLimit(clientIP) {
|
||||
t.Error("First request should pass")
|
||||
}
|
||||
|
||||
// Second request - should pass
|
||||
if !securityExt.checkRateLimit(clientIP) {
|
||||
t.Error("Second request should pass")
|
||||
}
|
||||
|
||||
// Third request - should be rate limited
|
||||
if securityExt.checkRateLimit(clientIP) {
|
||||
t.Error("Third request should be rate limited")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_GetMetrics(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{
|
||||
"blocked_ips": []interface{}{"192.168.1.1"},
|
||||
"allowed_ips": []interface{}{"192.168.1.2"},
|
||||
}, logger)
|
||||
|
||||
securityExt := ext.(*SecurityExtension)
|
||||
metrics := securityExt.GetMetrics()
|
||||
|
||||
if metrics["blocked_ips"].(int) != 1 {
|
||||
t.Errorf("Expected 1 blocked IP, got %v", metrics["blocked_ips"])
|
||||
}
|
||||
|
||||
if metrics["allowed_ips"].(int) != 1 {
|
||||
t.Errorf("Expected 1 allowed IP, got %v", metrics["allowed_ips"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecurityExtension_AddRemoveIPs(t *testing.T) {
|
||||
logger := newTestLogger()
|
||||
|
||||
ext, _ := NewSecurityExtension(map[string]interface{}{}, logger)
|
||||
securityExt := ext.(*SecurityExtension)
|
||||
|
||||
// Add blocked IP
|
||||
securityExt.AddBlockedIP("192.168.1.100")
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.RemoteAddr = "192.168.1.100:12345"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
handled, _ := ext.ProcessRequest(context.Background(), rr, req)
|
||||
if !handled {
|
||||
t.Error("Expected dynamically blocked IP to be blocked")
|
||||
}
|
||||
|
||||
// Remove blocked IP
|
||||
securityExt.RemoveBlockedIP("192.168.1.100")
|
||||
|
||||
rr = httptest.NewRecorder()
|
||||
handled, _ = ext.ProcessRequest(context.Background(), rr, req)
|
||||
if handled {
|
||||
t.Error("Expected removed blocked IP to pass through")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
// Package logging provides structured logging with zap
|
||||
package logging
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"gopkg.in/natefinch/lumberjack.v2"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/config"
|
||||
)
|
||||
|
||||
// Config is a simple configuration for basic logger setup
|
||||
type Config struct {
|
||||
Level string
|
||||
TimestampFormat string
|
||||
}
|
||||
|
||||
// Logger wraps zap.SugaredLogger with additional functionality
|
||||
type Logger struct {
|
||||
*zap.SugaredLogger
|
||||
zap *zap.Logger
|
||||
config *config.LoggingConfig
|
||||
name string
|
||||
}
|
||||
|
||||
// New creates a new Logger with basic configuration
|
||||
func New(cfg Config) (*Logger, error) {
|
||||
level := parseLevel(cfg.Level)
|
||||
|
||||
timestampFormat := cfg.TimestampFormat
|
||||
if timestampFormat == "" {
|
||||
timestampFormat = "2006-01-02 15:04:05"
|
||||
}
|
||||
|
||||
encoderConfig := zap.NewProductionEncoderConfig()
|
||||
encoderConfig.TimeKey = "timestamp"
|
||||
encoderConfig.EncodeTime = zapcore.TimeEncoderOfLayout(timestampFormat)
|
||||
encoderConfig.EncodeLevel = zapcore.CapitalColorLevelEncoder
|
||||
|
||||
core := zapcore.NewCore(
|
||||
zapcore.NewConsoleEncoder(encoderConfig),
|
||||
zapcore.AddSync(os.Stdout),
|
||||
level,
|
||||
)
|
||||
|
||||
zapLogger := zap.New(core)
|
||||
return &Logger{
|
||||
SugaredLogger: zapLogger.Sugar(),
|
||||
zap: zapLogger,
|
||||
name: "konduktor",
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NewFromConfig creates a Logger from full LoggingConfig
|
||||
func NewFromConfig(cfg config.LoggingConfig) (*Logger, error) {
|
||||
var cores []zapcore.Core
|
||||
|
||||
// Parse main level
|
||||
mainLevel := parseLevel(cfg.Level)
|
||||
|
||||
// Add console core if enabled
|
||||
if cfg.ConsoleOutput {
|
||||
consoleLevel := mainLevel
|
||||
if cfg.Console != nil && cfg.Console.Level != "" {
|
||||
consoleLevel = parseLevel(cfg.Console.Level)
|
||||
}
|
||||
|
||||
var consoleEncoder zapcore.Encoder
|
||||
formatConfig := cfg.Format
|
||||
if cfg.Console != nil {
|
||||
formatConfig = mergeFormatConfig(cfg.Format, cfg.Console.Format)
|
||||
}
|
||||
|
||||
encoderCfg := createEncoderConfig(formatConfig)
|
||||
if formatConfig.Type == "json" {
|
||||
consoleEncoder = zapcore.NewJSONEncoder(encoderCfg)
|
||||
} else {
|
||||
if formatConfig.UseColors {
|
||||
encoderCfg.EncodeLevel = zapcore.CapitalColorLevelEncoder
|
||||
}
|
||||
consoleEncoder = zapcore.NewConsoleEncoder(encoderCfg)
|
||||
}
|
||||
|
||||
consoleSyncer := zapcore.AddSync(os.Stdout)
|
||||
cores = append(cores, zapcore.NewCore(consoleEncoder, consoleSyncer, consoleLevel))
|
||||
}
|
||||
|
||||
// Add file cores
|
||||
for _, fileConfig := range cfg.Files {
|
||||
fileCore, err := createFileCore(fileConfig, cfg.Format, mainLevel)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create file logger for %s: %w", fileConfig.Path, err)
|
||||
}
|
||||
|
||||
// If specific loggers are configured, wrap with filter
|
||||
if len(fileConfig.Loggers) > 0 {
|
||||
fileCore = &filteredCore{
|
||||
Core: fileCore,
|
||||
loggers: fileConfig.Loggers,
|
||||
}
|
||||
}
|
||||
|
||||
cores = append(cores, fileCore)
|
||||
}
|
||||
|
||||
// If no cores configured, add default console
|
||||
if len(cores) == 0 {
|
||||
encoderCfg := zap.NewProductionEncoderConfig()
|
||||
encoderCfg.EncodeTime = zapcore.TimeEncoderOfLayout("2006-01-02 15:04:05")
|
||||
encoderCfg.EncodeLevel = zapcore.CapitalColorLevelEncoder
|
||||
cores = append(cores, zapcore.NewCore(
|
||||
zapcore.NewConsoleEncoder(encoderCfg),
|
||||
zapcore.AddSync(os.Stdout),
|
||||
mainLevel,
|
||||
))
|
||||
}
|
||||
|
||||
// Combine all cores
|
||||
core := zapcore.NewTee(cores...)
|
||||
zapLogger := zap.New(core, zap.AddCaller(), zap.AddCallerSkip(1))
|
||||
|
||||
return &Logger{
|
||||
SugaredLogger: zapLogger.Sugar(),
|
||||
zap: zapLogger,
|
||||
config: &cfg,
|
||||
name: "konduktor",
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Named returns a logger with a specific name (for filtering)
|
||||
func (l *Logger) Named(name string) *Logger {
|
||||
return &Logger{
|
||||
SugaredLogger: l.SugaredLogger.Named(name),
|
||||
zap: l.zap.Named(name),
|
||||
config: l.config,
|
||||
name: name,
|
||||
}
|
||||
}
|
||||
|
||||
// With returns a logger with additional fields
|
||||
func (l *Logger) With(args ...interface{}) *Logger {
|
||||
return &Logger{
|
||||
SugaredLogger: l.SugaredLogger.With(args...),
|
||||
zap: l.zap.Sugar().With(args...).Desugar(),
|
||||
config: l.config,
|
||||
name: l.name,
|
||||
}
|
||||
}
|
||||
|
||||
// Sync flushes any buffered log entries
|
||||
func (l *Logger) Sync() error {
|
||||
return l.zap.Sync()
|
||||
}
|
||||
|
||||
// GetZap returns the underlying zap.Logger
|
||||
func (l *Logger) GetZap() *zap.Logger {
|
||||
return l.zap
|
||||
}
|
||||
|
||||
// Debug logs a debug message
|
||||
func (l *Logger) Debug(msg string, keysAndValues ...interface{}) {
|
||||
l.SugaredLogger.Debugw(msg, keysAndValues...)
|
||||
}
|
||||
|
||||
// Info logs an info message
|
||||
func (l *Logger) Info(msg string, keysAndValues ...interface{}) {
|
||||
l.SugaredLogger.Infow(msg, keysAndValues...)
|
||||
}
|
||||
|
||||
// Warn logs a warning message
|
||||
func (l *Logger) Warn(msg string, keysAndValues ...interface{}) {
|
||||
l.SugaredLogger.Warnw(msg, keysAndValues...)
|
||||
}
|
||||
|
||||
// Error logs an error message
|
||||
func (l *Logger) Error(msg string, keysAndValues ...interface{}) {
|
||||
l.SugaredLogger.Errorw(msg, keysAndValues...)
|
||||
}
|
||||
|
||||
// Fatal logs a fatal message and exits
|
||||
func (l *Logger) Fatal(msg string, keysAndValues ...interface{}) {
|
||||
l.SugaredLogger.Fatalw(msg, keysAndValues...)
|
||||
}
|
||||
|
||||
// --- Helper functions ---
|
||||
|
||||
func parseLevel(level string) zapcore.Level {
|
||||
switch strings.ToUpper(level) {
|
||||
case "DEBUG":
|
||||
return zapcore.DebugLevel
|
||||
case "INFO":
|
||||
return zapcore.InfoLevel
|
||||
case "WARN", "WARNING":
|
||||
return zapcore.WarnLevel
|
||||
case "ERROR":
|
||||
return zapcore.ErrorLevel
|
||||
case "CRITICAL", "FATAL":
|
||||
return zapcore.FatalLevel
|
||||
default:
|
||||
return zapcore.InfoLevel
|
||||
}
|
||||
}
|
||||
|
||||
func createEncoderConfig(format config.LogFormatConfig) zapcore.EncoderConfig {
|
||||
timestampFormat := format.TimestampFormat
|
||||
if timestampFormat == "" {
|
||||
timestampFormat = "2006-01-02 15:04:05"
|
||||
}
|
||||
|
||||
cfg := zapcore.EncoderConfig{
|
||||
TimeKey: "timestamp",
|
||||
LevelKey: "level",
|
||||
NameKey: "logger",
|
||||
CallerKey: "caller",
|
||||
FunctionKey: zapcore.OmitKey,
|
||||
MessageKey: "msg",
|
||||
StacktraceKey: "stacktrace",
|
||||
LineEnding: zapcore.DefaultLineEnding,
|
||||
EncodeLevel: zapcore.CapitalLevelEncoder,
|
||||
EncodeTime: zapcore.TimeEncoderOfLayout(timestampFormat),
|
||||
EncodeDuration: zapcore.SecondsDurationEncoder,
|
||||
EncodeCaller: zapcore.ShortCallerEncoder,
|
||||
}
|
||||
|
||||
if !format.ShowModule {
|
||||
cfg.NameKey = zapcore.OmitKey
|
||||
}
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
func mergeFormatConfig(base, override config.LogFormatConfig) config.LogFormatConfig {
|
||||
result := base
|
||||
if override.Type != "" {
|
||||
result.Type = override.Type
|
||||
}
|
||||
if override.TimestampFormat != "" {
|
||||
result.TimestampFormat = override.TimestampFormat
|
||||
}
|
||||
// UseColors and ShowModule are bool - check if override has non-default
|
||||
result.UseColors = override.UseColors
|
||||
result.ShowModule = override.ShowModule
|
||||
return result
|
||||
}
|
||||
|
||||
func createFileCore(fileConfig config.FileLogConfig, defaultFormat config.LogFormatConfig, defaultLevel zapcore.Level) (zapcore.Core, error) {
|
||||
// Ensure directory exists
|
||||
dir := filepath.Dir(fileConfig.Path)
|
||||
if dir != "" && dir != "." {
|
||||
if err := os.MkdirAll(dir, 0755); err != nil {
|
||||
return nil, fmt.Errorf("failed to create log directory %s: %w", dir, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Configure log rotation with lumberjack
|
||||
maxSize := 10 // MB
|
||||
if fileConfig.MaxBytes > 0 {
|
||||
maxSize = int(fileConfig.MaxBytes / (1024 * 1024))
|
||||
if maxSize < 1 {
|
||||
maxSize = 1
|
||||
}
|
||||
}
|
||||
|
||||
backupCount := 5
|
||||
if fileConfig.BackupCount > 0 {
|
||||
backupCount = fileConfig.BackupCount
|
||||
}
|
||||
|
||||
rotator := &lumberjack.Logger{
|
||||
Filename: fileConfig.Path,
|
||||
MaxSize: maxSize,
|
||||
MaxBackups: backupCount,
|
||||
MaxAge: 30, // days
|
||||
Compress: true,
|
||||
}
|
||||
|
||||
// Determine level
|
||||
level := defaultLevel
|
||||
if fileConfig.Level != "" {
|
||||
level = parseLevel(fileConfig.Level)
|
||||
}
|
||||
|
||||
// Create encoder
|
||||
format := defaultFormat
|
||||
if fileConfig.Format.Type != "" {
|
||||
format = mergeFormatConfig(defaultFormat, fileConfig.Format)
|
||||
}
|
||||
// Files should not use colors
|
||||
format.UseColors = false
|
||||
|
||||
encoderConfig := createEncoderConfig(format)
|
||||
var encoder zapcore.Encoder
|
||||
if format.Type == "json" {
|
||||
encoder = zapcore.NewJSONEncoder(encoderConfig)
|
||||
} else {
|
||||
encoder = zapcore.NewConsoleEncoder(encoderConfig)
|
||||
}
|
||||
|
||||
return zapcore.NewCore(encoder, zapcore.AddSync(rotator), level), nil
|
||||
}
|
||||
|
||||
// filteredCore wraps a Core to filter by logger name
|
||||
type filteredCore struct {
|
||||
zapcore.Core
|
||||
loggers []string
|
||||
}
|
||||
|
||||
func (c *filteredCore) Check(entry zapcore.Entry, ce *zapcore.CheckedEntry) *zapcore.CheckedEntry {
|
||||
if !c.shouldLog(entry.LoggerName) {
|
||||
return ce
|
||||
}
|
||||
return c.Core.Check(entry, ce)
|
||||
}
|
||||
|
||||
func (c *filteredCore) shouldLog(loggerName string) bool {
|
||||
if len(c.loggers) == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
for _, allowed := range c.loggers {
|
||||
if loggerName == allowed || strings.HasPrefix(loggerName, allowed+".") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (c *filteredCore) With(fields []zapcore.Field) zapcore.Core {
|
||||
return &filteredCore{
|
||||
Core: c.Core.With(fields),
|
||||
loggers: c.loggers,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,209 @@
|
||||
package logging
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/config"
|
||||
)
|
||||
|
||||
func TestNew(t *testing.T) {
|
||||
logger, err := New(Config{Level: "INFO"})
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if logger == nil {
|
||||
t.Fatal("Expected logger, got nil")
|
||||
}
|
||||
|
||||
if logger.name != "konduktor" {
|
||||
t.Errorf("Expected name konduktor, got %s", logger.name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_DefaultTimestampFormat(t *testing.T) {
|
||||
logger, err := New(Config{Level: "DEBUG"})
|
||||
if err != nil {
|
||||
t.Fatalf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
// Logger should be created successfully
|
||||
if logger == nil {
|
||||
t.Fatal("Expected logger, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_CustomTimestampFormat(t *testing.T) {
|
||||
logger, err := New(Config{
|
||||
Level: "DEBUG",
|
||||
TimestampFormat: "15:04:05",
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if logger == nil {
|
||||
t.Fatal("Expected logger, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewFromConfig(t *testing.T) {
|
||||
cfg := config.LoggingConfig{
|
||||
Level: "DEBUG",
|
||||
ConsoleOutput: true,
|
||||
Format: config.LogFormatConfig{
|
||||
Type: "standard",
|
||||
UseColors: true,
|
||||
ShowModule: true,
|
||||
TimestampFormat: "2006-01-02 15:04:05",
|
||||
},
|
||||
}
|
||||
|
||||
logger, err := NewFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if logger == nil {
|
||||
t.Fatal("Expected logger, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewFromConfig_WithConsole(t *testing.T) {
|
||||
cfg := config.LoggingConfig{
|
||||
Level: "INFO",
|
||||
ConsoleOutput: true,
|
||||
Format: config.LogFormatConfig{
|
||||
Type: "standard",
|
||||
UseColors: true,
|
||||
},
|
||||
Console: &config.ConsoleLogConfig{
|
||||
Level: "DEBUG",
|
||||
Format: config.LogFormatConfig{
|
||||
Type: "standard",
|
||||
UseColors: false,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
logger, err := NewFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if logger == nil {
|
||||
t.Fatal("Expected logger, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogger_Debug(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "DEBUG"})
|
||||
|
||||
// Should not panic
|
||||
logger.Debug("test message", "key", "value")
|
||||
}
|
||||
|
||||
func TestLogger_Info(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "INFO"})
|
||||
|
||||
// Should not panic
|
||||
logger.Info("test message", "key", "value")
|
||||
}
|
||||
|
||||
func TestLogger_Warn(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "WARN"})
|
||||
|
||||
// Should not panic
|
||||
logger.Warn("test message", "key", "value")
|
||||
}
|
||||
|
||||
func TestLogger_Error(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "ERROR"})
|
||||
|
||||
// Should not panic
|
||||
logger.Error("test message", "key", "value")
|
||||
}
|
||||
|
||||
func TestLogger_Named(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "INFO"})
|
||||
named := logger.Named("test.module")
|
||||
|
||||
if named == nil {
|
||||
t.Fatal("Expected named logger, got nil")
|
||||
}
|
||||
|
||||
if named.name != "test.module" {
|
||||
t.Errorf("Expected name 'test.module', got %s", named.name)
|
||||
}
|
||||
|
||||
// Should not panic
|
||||
named.Info("test from named logger")
|
||||
}
|
||||
|
||||
func TestLogger_With(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "INFO"})
|
||||
withFields := logger.With("service", "test")
|
||||
|
||||
if withFields == nil {
|
||||
t.Fatal("Expected logger with fields, got nil")
|
||||
}
|
||||
|
||||
// Should not panic
|
||||
withFields.Info("test with fields")
|
||||
}
|
||||
|
||||
func TestLogger_Sync(t *testing.T) {
|
||||
logger, _ := New(Config{Level: "INFO"})
|
||||
|
||||
// Should not panic
|
||||
err := logger.Sync()
|
||||
// Sync may return an error for stdout on some systems, ignore it
|
||||
_ = err
|
||||
}
|
||||
|
||||
func TestParseLevel(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"DEBUG", "debug"},
|
||||
{"INFO", "info"},
|
||||
{"WARN", "warn"},
|
||||
{"WARNING", "warn"},
|
||||
{"ERROR", "error"},
|
||||
{"CRITICAL", "fatal"},
|
||||
{"FATAL", "fatal"},
|
||||
{"invalid", "info"}, // defaults to INFO
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
level := parseLevel(tt.input)
|
||||
if level.String() != tt.expected {
|
||||
t.Errorf("parseLevel(%s) = %s, want %s", tt.input, level.String(), tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Benchmarks ==============
|
||||
|
||||
func BenchmarkLogger_Info(b *testing.B) {
|
||||
logger, _ := New(Config{Level: "INFO"})
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
logger.Info("test message", "key", "value")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLogger_Debug_Filtered(b *testing.B) {
|
||||
logger, _ := New(Config{Level: "ERROR"})
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
logger.Debug("test message", "key", "value")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
type responseWriter struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
size int
|
||||
}
|
||||
|
||||
func (rw *responseWriter) WriteHeader(code int) {
|
||||
rw.status = code
|
||||
rw.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
func (rw *responseWriter) Write(b []byte) (int, error) {
|
||||
size, err := rw.ResponseWriter.Write(b)
|
||||
rw.size += size
|
||||
return size, err
|
||||
}
|
||||
|
||||
func ServerHeader(next http.Handler, version string) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Server", fmt.Sprintf("konduktor/%s", version))
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func AccessLog(next http.Handler, logger *logging.Logger) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
start := time.Now()
|
||||
|
||||
wrapped := &responseWriter{
|
||||
ResponseWriter: w,
|
||||
status: http.StatusOK,
|
||||
}
|
||||
|
||||
next.ServeHTTP(wrapped, r)
|
||||
|
||||
duration := time.Since(start)
|
||||
|
||||
logger.Info("HTTP request",
|
||||
"method", r.Method,
|
||||
"path", r.URL.Path,
|
||||
"status", wrapped.status,
|
||||
"duration_ms", duration.Milliseconds(),
|
||||
"client_ip", r.RemoteAddr,
|
||||
"user_agent", r.UserAgent(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
func Recovery(next http.Handler, logger *logging.Logger) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
logger.Error("Panic recovered",
|
||||
"error", fmt.Sprintf("%v", err),
|
||||
"stack", string(debug.Stack()),
|
||||
)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}()
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
// ============== ServerHeader Tests ==============
|
||||
|
||||
func TestServerHeader(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
|
||||
wrapped := ServerHeader(handler, "1.0.0")
|
||||
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
serverHeader := rr.Header().Get("Server")
|
||||
if serverHeader != "konduktor/1.0.0" {
|
||||
t.Errorf("Expected Server header 'konduktor/1.0.0', got '%s'", serverHeader)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== AccessLog Tests ==============
|
||||
|
||||
func TestAccessLog(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte("Hello"))
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "INFO"})
|
||||
wrapped := AccessLog(handler, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccessLog_CapturesStatusCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
}{
|
||||
{"OK", http.StatusOK},
|
||||
{"NotFound", http.StatusNotFound},
|
||||
{"InternalError", http.StatusInternalServerError},
|
||||
{"Redirect", http.StatusMovedPermanently},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(tt.statusCode)
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "INFO"})
|
||||
wrapped := AccessLog(handler, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != tt.statusCode {
|
||||
t.Errorf("Expected status %d, got %d", tt.statusCode, rr.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Recovery Tests ==============
|
||||
|
||||
func TestRecovery_NoPanic(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "INFO"})
|
||||
wrapped := Recovery(handler, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
|
||||
if rr.Body.String() != "OK" {
|
||||
t.Errorf("Expected body 'OK', got '%s'", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecovery_WithPanic(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
panic("test panic")
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "ERROR"})
|
||||
wrapped := Recovery(handler, logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
// Should not panic
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusInternalServerError {
|
||||
t.Errorf("Expected status 500, got %d", rr.Code)
|
||||
}
|
||||
|
||||
if !strings.Contains(rr.Body.String(), "Internal Server Error") {
|
||||
t.Errorf("Expected 'Internal Server Error' in body, got '%s'", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== responseWriter Tests ==============
|
||||
|
||||
func TestResponseWriter_WriteHeader(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
rw := &responseWriter{ResponseWriter: rr, status: http.StatusOK}
|
||||
|
||||
rw.WriteHeader(http.StatusNotFound)
|
||||
|
||||
if rw.status != http.StatusNotFound {
|
||||
t.Errorf("Expected status 404, got %d", rw.status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponseWriter_Write(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
rw := &responseWriter{ResponseWriter: rr, status: http.StatusOK}
|
||||
|
||||
n, err := rw.Write([]byte("Hello World"))
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if n != 11 {
|
||||
t.Errorf("Expected 11 bytes written, got %d", n)
|
||||
}
|
||||
|
||||
if rw.size != 11 {
|
||||
t.Errorf("Expected size 11, got %d", rw.size)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Middleware Chain Tests ==============
|
||||
|
||||
func TestMiddlewareChain(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "INFO"})
|
||||
|
||||
// Apply middleware chain
|
||||
wrapped := Recovery(AccessLog(ServerHeader(handler, "1.0.0"), logger), logger)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
|
||||
// Check all middleware worked
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
|
||||
if rr.Header().Get("Server") != "konduktor/1.0.0" {
|
||||
t.Errorf("Expected Server header")
|
||||
}
|
||||
|
||||
if rr.Body.String() != "OK" {
|
||||
t.Errorf("Expected body 'OK', got '%s'", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Benchmarks ==============
|
||||
|
||||
func BenchmarkServerHeader(b *testing.B) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
|
||||
wrapped := ServerHeader(handler, "1.0.0")
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
rr := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkAccessLog(b *testing.B) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "ERROR"}) // Minimize logging overhead
|
||||
wrapped := AccessLog(handler, logger)
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
rr := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkRecovery(b *testing.B) {
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
|
||||
logger, _ := logging.New(logging.Config{Level: "ERROR"})
|
||||
wrapped := Recovery(handler, logger)
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
rr := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(rr, req)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package pathmatcher
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type MountedPath struct {
|
||||
path string
|
||||
name string
|
||||
stripPath bool
|
||||
}
|
||||
|
||||
func NewMountedPath(path string, opts ...MountedPathOption) *MountedPath {
|
||||
// Normalize: remove trailing slash (except for root)
|
||||
normalizedPath := strings.TrimSuffix(path, "/")
|
||||
if normalizedPath == "" {
|
||||
normalizedPath = ""
|
||||
}
|
||||
|
||||
m := &MountedPath{
|
||||
path: normalizedPath,
|
||||
name: normalizedPath,
|
||||
stripPath: true,
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(m)
|
||||
}
|
||||
|
||||
if m.name == "" {
|
||||
m.name = normalizedPath
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
type MountedPathOption func(*MountedPath)
|
||||
|
||||
func WithName(name string) MountedPathOption {
|
||||
return func(m *MountedPath) {
|
||||
m.name = name
|
||||
}
|
||||
}
|
||||
|
||||
func WithStripPath(strip bool) MountedPathOption {
|
||||
return func(m *MountedPath) {
|
||||
m.stripPath = strip
|
||||
}
|
||||
}
|
||||
|
||||
func (m *MountedPath) Path() string {
|
||||
return m.path
|
||||
}
|
||||
|
||||
func (m *MountedPath) Name() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MountedPath) StripPath() bool {
|
||||
return m.stripPath
|
||||
}
|
||||
|
||||
func (m *MountedPath) Matches(requestPath string) bool {
|
||||
// Empty or "/" mount matches everything
|
||||
if m.path == "" || m.path == "/" {
|
||||
return true
|
||||
}
|
||||
|
||||
// Request path must be at least as long as mount path
|
||||
if len(requestPath) < len(m.path) {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check if request path starts with mount path
|
||||
if !strings.HasPrefix(requestPath, m.path) {
|
||||
return false
|
||||
}
|
||||
|
||||
// If paths are equal length, it's a match
|
||||
if len(requestPath) == len(m.path) {
|
||||
return true
|
||||
}
|
||||
|
||||
// Otherwise, next char must be '/' to prevent /api matching /api-v2
|
||||
return requestPath[len(m.path)] == '/'
|
||||
}
|
||||
|
||||
func (m *MountedPath) GetModifiedPath(requestPath string) string {
|
||||
if !m.stripPath {
|
||||
return requestPath
|
||||
}
|
||||
|
||||
// Root mount doesn't strip anything
|
||||
if m.path == "" || m.path == "/" {
|
||||
return requestPath
|
||||
}
|
||||
|
||||
// Strip the prefix
|
||||
modified := strings.TrimPrefix(requestPath, m.path)
|
||||
|
||||
// Ensure result starts with /
|
||||
if modified == "" || modified[0] != '/' {
|
||||
modified = "/" + modified
|
||||
}
|
||||
|
||||
return modified
|
||||
}
|
||||
|
||||
type MountManager struct {
|
||||
mounts []*MountedPath
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func NewMountManager() *MountManager {
|
||||
return &MountManager{
|
||||
mounts: make([]*MountedPath, 0),
|
||||
}
|
||||
}
|
||||
|
||||
func (mm *MountManager) AddMount(mount *MountedPath) {
|
||||
mm.mu.Lock()
|
||||
defer mm.mu.Unlock()
|
||||
|
||||
// Insert in sorted order (longer paths first)
|
||||
inserted := false
|
||||
for i, existing := range mm.mounts {
|
||||
if len(mount.path) > len(existing.path) {
|
||||
// Insert at position i
|
||||
mm.mounts = append(mm.mounts[:i], append([]*MountedPath{mount}, mm.mounts[i:]...)...)
|
||||
inserted = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !inserted {
|
||||
mm.mounts = append(mm.mounts, mount)
|
||||
}
|
||||
}
|
||||
|
||||
func (mm *MountManager) RemoveMount(path string) bool {
|
||||
mm.mu.Lock()
|
||||
defer mm.mu.Unlock()
|
||||
|
||||
normalizedPath := strings.TrimSuffix(path, "/")
|
||||
|
||||
for i, mount := range mm.mounts {
|
||||
if mount.path == normalizedPath {
|
||||
mm.mounts = append(mm.mounts[:i], mm.mounts[i+1:]...)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (mm *MountManager) GetMount(requestPath string) *MountedPath {
|
||||
mm.mu.RLock()
|
||||
defer mm.mu.RUnlock()
|
||||
|
||||
// Mounts are sorted by path length (longest first)
|
||||
// so the first match is the best match
|
||||
for _, mount := range mm.mounts {
|
||||
if mount.Matches(requestPath) {
|
||||
return mount
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (mm *MountManager) MountCount() int {
|
||||
mm.mu.RLock()
|
||||
defer mm.mu.RUnlock()
|
||||
return len(mm.mounts)
|
||||
}
|
||||
|
||||
func (mm *MountManager) Mounts() []*MountedPath {
|
||||
mm.mu.RLock()
|
||||
defer mm.mu.RUnlock()
|
||||
|
||||
result := make([]*MountedPath, len(mm.mounts))
|
||||
copy(result, mm.mounts)
|
||||
return result
|
||||
}
|
||||
|
||||
func (mm *MountManager) ListMounts() []map[string]interface{} {
|
||||
mm.mu.RLock()
|
||||
defer mm.mu.RUnlock()
|
||||
|
||||
result := make([]map[string]interface{}, len(mm.mounts))
|
||||
for i, mount := range mm.mounts {
|
||||
result[i] = map[string]interface{}{
|
||||
"path": mount.path,
|
||||
"name": mount.name,
|
||||
"strip_path": mount.stripPath,
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Utility functions
|
||||
|
||||
func PathMatchesPrefix(requestPath, prefix string) bool {
|
||||
// Normalize prefix
|
||||
prefix = strings.TrimSuffix(prefix, "/")
|
||||
|
||||
// Empty or "/" prefix matches everything
|
||||
if prefix == "" || prefix == "/" {
|
||||
return true
|
||||
}
|
||||
|
||||
// Request path must be at least as long as prefix
|
||||
if len(requestPath) < len(prefix) {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check if request path starts with prefix
|
||||
if !strings.HasPrefix(requestPath, prefix) {
|
||||
return false
|
||||
}
|
||||
|
||||
// If paths are equal length, it's a match
|
||||
if len(requestPath) == len(prefix) {
|
||||
return true
|
||||
}
|
||||
|
||||
// Otherwise, next char must be '/'
|
||||
return requestPath[len(prefix)] == '/'
|
||||
}
|
||||
|
||||
func StripPathPrefix(requestPath, prefix string) string {
|
||||
// Normalize prefix
|
||||
prefix = strings.TrimSuffix(prefix, "/")
|
||||
|
||||
// Empty or "/" prefix doesn't strip anything
|
||||
if prefix == "" || prefix == "/" {
|
||||
return requestPath
|
||||
}
|
||||
|
||||
// Strip the prefix
|
||||
modified := strings.TrimPrefix(requestPath, prefix)
|
||||
|
||||
// Ensure result starts with /
|
||||
if modified == "" || modified[0] != '/' {
|
||||
modified = "/" + modified
|
||||
}
|
||||
|
||||
return modified
|
||||
}
|
||||
|
||||
func MatchAndModifyPath(requestPath, prefix string, stripPath bool) (matches bool, modifiedPath string) {
|
||||
if !PathMatchesPrefix(requestPath, prefix) {
|
||||
return false, ""
|
||||
}
|
||||
|
||||
if stripPath {
|
||||
return true, StripPathPrefix(requestPath, prefix)
|
||||
}
|
||||
|
||||
return true, requestPath
|
||||
}
|
||||
@@ -0,0 +1,460 @@
|
||||
package pathmatcher
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ============== MountedPath Tests ==============
|
||||
|
||||
func TestMountedPath_RootMountMatchesEverything(t *testing.T) {
|
||||
mount := NewMountedPath("")
|
||||
|
||||
tests := []string{"/", "/api", "/api/users", "/anything/at/all"}
|
||||
|
||||
for _, path := range tests {
|
||||
if !mount.Matches(path) {
|
||||
t.Errorf("Root mount should match %s", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_SlashRootMountMatchesEverything(t *testing.T) {
|
||||
mount := NewMountedPath("/")
|
||||
|
||||
tests := []string{"/", "/api", "/api/users"}
|
||||
|
||||
for _, path := range tests {
|
||||
if !mount.Matches(path) {
|
||||
t.Errorf("'/' mount should match %s", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_ExactPathMatch(t *testing.T) {
|
||||
mount := NewMountedPath("/api")
|
||||
|
||||
tests := []struct {
|
||||
path string
|
||||
expected bool
|
||||
}{
|
||||
{"/api", true},
|
||||
{"/api/", true},
|
||||
{"/api/users", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := mount.Matches(tt.path); got != tt.expected {
|
||||
t.Errorf("Matches(%s) = %v, want %v", tt.path, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_NoFalsePrefixMatch(t *testing.T) {
|
||||
mount := NewMountedPath("/api")
|
||||
|
||||
tests := []string{"/api-v2", "/api2", "/apiv2"}
|
||||
|
||||
for _, path := range tests {
|
||||
if mount.Matches(path) {
|
||||
t.Errorf("/api should not match %s", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_ShorterPathNoMatch(t *testing.T) {
|
||||
mount := NewMountedPath("/api/v1")
|
||||
|
||||
tests := []string{"/api", "/ap", "/"}
|
||||
|
||||
for _, path := range tests {
|
||||
if mount.Matches(path) {
|
||||
t.Errorf("/api/v1 should not match shorter path %s", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_TrailingSlashNormalized(t *testing.T) {
|
||||
mount1 := NewMountedPath("/api/")
|
||||
mount2 := NewMountedPath("/api")
|
||||
|
||||
if mount1.Path() != "/api" {
|
||||
t.Errorf("Expected path /api, got %s", mount1.Path())
|
||||
}
|
||||
|
||||
if mount2.Path() != "/api" {
|
||||
t.Errorf("Expected path /api, got %s", mount2.Path())
|
||||
}
|
||||
|
||||
if !mount1.Matches("/api/users") {
|
||||
t.Error("mount1 should match /api/users")
|
||||
}
|
||||
|
||||
if !mount2.Matches("/api/users") {
|
||||
t.Error("mount2 should match /api/users")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_GetModifiedPathStripsPrefix(t *testing.T) {
|
||||
mount := NewMountedPath("/api")
|
||||
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"/api", "/"},
|
||||
{"/api/", "/"},
|
||||
{"/api/users", "/users"},
|
||||
{"/api/users/123", "/users/123"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := mount.GetModifiedPath(tt.input); got != tt.expected {
|
||||
t.Errorf("GetModifiedPath(%s) = %s, want %s", tt.input, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_GetModifiedPathNoStrip(t *testing.T) {
|
||||
mount := NewMountedPath("/api", WithStripPath(false))
|
||||
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"/api/users", "/api/users"},
|
||||
{"/api", "/api"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := mount.GetModifiedPath(tt.input); got != tt.expected {
|
||||
t.Errorf("GetModifiedPath(%s) = %s, want %s", tt.input, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_RootMountModifiedPath(t *testing.T) {
|
||||
mount := NewMountedPath("")
|
||||
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"/api/users", "/api/users"},
|
||||
{"/", "/"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := mount.GetModifiedPath(tt.input); got != tt.expected {
|
||||
t.Errorf("GetModifiedPath(%s) = %s, want %s", tt.input, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedPath_NameProperty(t *testing.T) {
|
||||
mount1 := NewMountedPath("/api")
|
||||
mount2 := NewMountedPath("/api", WithName("API Mount"))
|
||||
|
||||
if mount1.Name() != "/api" {
|
||||
t.Errorf("Expected name /api, got %s", mount1.Name())
|
||||
}
|
||||
|
||||
if mount2.Name() != "API Mount" {
|
||||
t.Errorf("Expected name 'API Mount', got %s", mount2.Name())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== MountManager Tests ==============
|
||||
|
||||
func TestMountManager_EmptyManager(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
if got := manager.GetMount("/api"); got != nil {
|
||||
t.Error("Empty manager should return nil")
|
||||
}
|
||||
|
||||
if got := manager.MountCount(); got != 0 {
|
||||
t.Errorf("Expected mount count 0, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_AddMount(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
mount := NewMountedPath("/api")
|
||||
|
||||
manager.AddMount(mount)
|
||||
|
||||
if manager.MountCount() != 1 {
|
||||
t.Errorf("Expected mount count 1, got %d", manager.MountCount())
|
||||
}
|
||||
|
||||
if got := manager.GetMount("/api/users"); got != mount {
|
||||
t.Error("GetMount should return the added mount")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_LongestPrefixMatching(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
apiMount := NewMountedPath("/api", WithName("api"))
|
||||
apiV1Mount := NewMountedPath("/api/v1", WithName("api_v1"))
|
||||
apiV2Mount := NewMountedPath("/api/v2", WithName("api_v2"))
|
||||
|
||||
manager.AddMount(apiMount)
|
||||
manager.AddMount(apiV2Mount)
|
||||
manager.AddMount(apiV1Mount)
|
||||
|
||||
tests := []struct {
|
||||
path string
|
||||
expectedName string
|
||||
}{
|
||||
{"/api/v1/users", "api_v1"},
|
||||
{"/api/v2/items", "api_v2"},
|
||||
{"/api/v3/other", "api"},
|
||||
{"/api", "api"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
got := manager.GetMount(tt.path)
|
||||
if got == nil {
|
||||
t.Errorf("GetMount(%s) returned nil, want mount with name %s", tt.path, tt.expectedName)
|
||||
continue
|
||||
}
|
||||
if got.Name() != tt.expectedName {
|
||||
t.Errorf("GetMount(%s).Name() = %s, want %s", tt.path, got.Name(), tt.expectedName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_RemoveMount(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
manager.AddMount(NewMountedPath("/api"))
|
||||
manager.AddMount(NewMountedPath("/admin"))
|
||||
|
||||
if manager.MountCount() != 2 {
|
||||
t.Errorf("Expected mount count 2, got %d", manager.MountCount())
|
||||
}
|
||||
|
||||
result := manager.RemoveMount("/api")
|
||||
|
||||
if !result {
|
||||
t.Error("RemoveMount should return true")
|
||||
}
|
||||
|
||||
if manager.MountCount() != 1 {
|
||||
t.Errorf("Expected mount count 1, got %d", manager.MountCount())
|
||||
}
|
||||
|
||||
if manager.GetMount("/api/users") != nil {
|
||||
t.Error("GetMount(/api/users) should return nil after removal")
|
||||
}
|
||||
|
||||
if manager.GetMount("/admin/users") == nil {
|
||||
t.Error("GetMount(/admin/users) should still work")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_RemoveNonexistentMount(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
result := manager.RemoveMount("/api")
|
||||
|
||||
if result {
|
||||
t.Error("RemoveMount should return false for nonexistent mount")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_ListMounts(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
manager.AddMount(NewMountedPath("/api", WithName("API")))
|
||||
manager.AddMount(NewMountedPath("/admin", WithName("Admin")))
|
||||
|
||||
mounts := manager.ListMounts()
|
||||
|
||||
if len(mounts) != 2 {
|
||||
t.Errorf("Expected 2 mounts, got %d", len(mounts))
|
||||
}
|
||||
|
||||
for _, m := range mounts {
|
||||
if _, ok := m["path"]; !ok {
|
||||
t.Error("Mount should have 'path' key")
|
||||
}
|
||||
if _, ok := m["name"]; !ok {
|
||||
t.Error("Mount should have 'name' key")
|
||||
}
|
||||
if _, ok := m["strip_path"]; !ok {
|
||||
t.Error("Mount should have 'strip_path' key")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountManager_MountsReturnsCopy(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
manager.AddMount(NewMountedPath("/api"))
|
||||
|
||||
mounts1 := manager.Mounts()
|
||||
mounts2 := manager.Mounts()
|
||||
|
||||
if &mounts1[0] == &mounts2[0] {
|
||||
t.Error("Mounts() should return different slices")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Utility Functions Tests ==============
|
||||
|
||||
func TestPathMatchesPrefix_Basic(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
prefix string
|
||||
expected bool
|
||||
}{
|
||||
{"/api/users", "/api", true},
|
||||
{"/api", "/api", true},
|
||||
{"/api-v2", "/api", false},
|
||||
{"/ap", "/api", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := PathMatchesPrefix(tt.path, tt.prefix); got != tt.expected {
|
||||
t.Errorf("PathMatchesPrefix(%s, %s) = %v, want %v", tt.path, tt.prefix, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathMatchesPrefix_Root(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
prefix string
|
||||
expected bool
|
||||
}{
|
||||
{"/anything", "", true},
|
||||
{"/anything", "/", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := PathMatchesPrefix(tt.path, tt.prefix); got != tt.expected {
|
||||
t.Errorf("PathMatchesPrefix(%s, %s) = %v, want %v", tt.path, tt.prefix, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripPathPrefix_Basic(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
prefix string
|
||||
expected string
|
||||
}{
|
||||
{"/api/users", "/api", "/users"},
|
||||
{"/api", "/api", "/"},
|
||||
{"/api/", "/api", "/"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := StripPathPrefix(tt.path, tt.prefix); got != tt.expected {
|
||||
t.Errorf("StripPathPrefix(%s, %s) = %s, want %s", tt.path, tt.prefix, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripPathPrefix_Root(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
prefix string
|
||||
expected string
|
||||
}{
|
||||
{"/api/users", "", "/api/users"},
|
||||
{"/api/users", "/", "/api/users"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := StripPathPrefix(tt.path, tt.prefix); got != tt.expected {
|
||||
t.Errorf("StripPathPrefix(%s, %s) = %s, want %s", tt.path, tt.prefix, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchAndModifyPath_Combined(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
prefix string
|
||||
stripPath bool
|
||||
wantMatches bool
|
||||
wantModified string
|
||||
}{
|
||||
{"/api/users", "/api", true, true, "/users"},
|
||||
{"/api", "/api", true, true, "/"},
|
||||
{"/other", "/api", true, false, ""},
|
||||
{"/api/users", "/api", false, true, "/api/users"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
matches, modified := MatchAndModifyPath(tt.path, tt.prefix, tt.stripPath)
|
||||
if matches != tt.wantMatches {
|
||||
t.Errorf("MatchAndModifyPath(%s, %s, %v) matches = %v, want %v",
|
||||
tt.path, tt.prefix, tt.stripPath, matches, tt.wantMatches)
|
||||
}
|
||||
if modified != tt.wantModified {
|
||||
t.Errorf("MatchAndModifyPath(%s, %s, %v) modified = %s, want %s",
|
||||
tt.path, tt.prefix, tt.stripPath, modified, tt.wantModified)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Performance Tests ==============
|
||||
|
||||
func TestPerformance_ManyMatches(t *testing.T) {
|
||||
mount := NewMountedPath("/api/v1/users")
|
||||
|
||||
for i := 0; i < 10000; i++ {
|
||||
if !mount.Matches("/api/v1/users/123/posts") {
|
||||
t.Fatal("Should match")
|
||||
}
|
||||
if mount.Matches("/other/path") {
|
||||
t.Fatal("Should not match")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerformance_ManyMounts(t *testing.T) {
|
||||
manager := NewMountManager()
|
||||
|
||||
for i := 0; i < 100; i++ {
|
||||
manager.AddMount(NewMountedPath("/api/v" + string(rune('0'+i%10)) + string(rune('0'+i/10))))
|
||||
}
|
||||
|
||||
if manager.MountCount() != 100 {
|
||||
t.Errorf("Expected 100 mounts, got %d", manager.MountCount())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Benchmarks ==============
|
||||
|
||||
func BenchmarkMountedPath_Matches(b *testing.B) {
|
||||
mount := NewMountedPath("/api/v1")
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
mount.Matches("/api/v1/users/123")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkMountManager_GetMount(b *testing.B) {
|
||||
manager := NewMountManager()
|
||||
|
||||
for i := 0; i < 20; i++ {
|
||||
manager.AddMount(NewMountedPath("/api/v" + string(rune('0'+i%10))))
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
manager.GetMount("/api/v5/users/123")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkPathMatchesPrefix(b *testing.B) {
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
PathMatchesPrefix("/api/v1/users/123", "/api/v1")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
// Package proxy provides reverse proxy functionality for Konduktor
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
// Target is the backend server URL
|
||||
Target string
|
||||
|
||||
// Timeout is the request timeout (default: 30s)
|
||||
Timeout time.Duration
|
||||
|
||||
// Headers are additional headers to add to requests
|
||||
Headers map[string]string
|
||||
|
||||
// StripPrefix removes this prefix from the request path
|
||||
StripPrefix string
|
||||
|
||||
// PreserveHost keeps the original Host header
|
||||
PreserveHost bool
|
||||
|
||||
// IgnoreRequestPath ignores the request path and uses only the target path
|
||||
// This is useful for exact match routes where target URL should be used as-is
|
||||
IgnoreRequestPath bool
|
||||
}
|
||||
|
||||
type ReverseProxy struct {
|
||||
config *Config
|
||||
targetURL *url.URL
|
||||
httpClient *http.Client
|
||||
logger *logging.Logger
|
||||
}
|
||||
|
||||
func New(cfg *Config, logger *logging.Logger) (*ReverseProxy, error) {
|
||||
if cfg.Target == "" {
|
||||
return nil, fmt.Errorf("proxy target is required")
|
||||
}
|
||||
|
||||
targetURL, err := url.Parse(cfg.Target)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid proxy target URL: %w", err)
|
||||
}
|
||||
|
||||
timeout := cfg.Timeout
|
||||
if timeout == 0 {
|
||||
timeout = 30 * time.Second
|
||||
}
|
||||
|
||||
transport := &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}).DialContext,
|
||||
MaxIdleConns: 100,
|
||||
MaxIdleConnsPerHost: 10,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
ResponseHeaderTimeout: timeout,
|
||||
}
|
||||
|
||||
return &ReverseProxy{
|
||||
config: cfg,
|
||||
targetURL: targetURL,
|
||||
httpClient: &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: timeout,
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse // Don't follow redirects
|
||||
},
|
||||
},
|
||||
logger: logger,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
rp.ProxyRequest(w, r, nil)
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) ProxyRequest(w http.ResponseWriter, r *http.Request, params map[string]string) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Build target URL
|
||||
targetURL := rp.buildTargetURL(r)
|
||||
|
||||
// Create proxy request
|
||||
proxyReq, err := rp.createProxyRequest(ctx, r, targetURL)
|
||||
if err != nil {
|
||||
rp.handleError(w, http.StatusInternalServerError, "Failed to create proxy request", err)
|
||||
return
|
||||
}
|
||||
|
||||
// Add custom headers with parameter substitution
|
||||
rp.addCustomHeaders(proxyReq, r, params)
|
||||
|
||||
// Execute request
|
||||
resp, err := rp.httpClient.Do(proxyReq)
|
||||
if err != nil {
|
||||
rp.handleProxyError(w, err)
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Copy response
|
||||
rp.copyResponse(w, resp)
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) buildTargetURL(r *http.Request) *url.URL {
|
||||
targetURL := *rp.targetURL
|
||||
|
||||
// If ignoring request path, use target URL path as-is
|
||||
if rp.config.IgnoreRequestPath {
|
||||
// Preserve query string only
|
||||
targetURL.RawQuery = r.URL.RawQuery
|
||||
return &targetURL
|
||||
}
|
||||
|
||||
// Strip prefix if configured
|
||||
path := r.URL.Path
|
||||
if rp.config.StripPrefix != "" {
|
||||
path = strings.TrimPrefix(path, rp.config.StripPrefix)
|
||||
if path == "" || path[0] != '/' {
|
||||
path = "/" + path
|
||||
}
|
||||
}
|
||||
|
||||
// If target URL has a non-empty path, combine it with the request path
|
||||
if rp.targetURL.Path != "" && rp.targetURL.Path != "/" {
|
||||
// Combine target path with request path
|
||||
targetURL.Path = strings.TrimSuffix(rp.targetURL.Path, "/") + path
|
||||
} else {
|
||||
// No path in target, use request path as-is
|
||||
targetURL.Path = path
|
||||
}
|
||||
|
||||
// Preserve query string
|
||||
targetURL.RawQuery = r.URL.RawQuery
|
||||
|
||||
return &targetURL
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) createProxyRequest(ctx context.Context, r *http.Request, targetURL *url.URL) (*http.Request, error) {
|
||||
proxyReq, err := http.NewRequestWithContext(ctx, r.Method, targetURL.String(), r.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Copy ContentLength
|
||||
proxyReq.ContentLength = r.ContentLength
|
||||
|
||||
// Copy headers
|
||||
for key, values := range r.Header {
|
||||
for _, value := range values {
|
||||
proxyReq.Header.Add(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// Set/update Host header
|
||||
if rp.config.PreserveHost {
|
||||
proxyReq.Host = r.Host
|
||||
} else {
|
||||
proxyReq.Host = targetURL.Host
|
||||
}
|
||||
|
||||
// Remove hop-by-hop headers
|
||||
removeHopByHopHeaders(proxyReq.Header)
|
||||
|
||||
return proxyReq, nil
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) addCustomHeaders(proxyReq *http.Request, originalReq *http.Request, params map[string]string) {
|
||||
// Add X-Forwarded headers
|
||||
clientIP := getClientIP(originalReq)
|
||||
if prior := originalReq.Header.Get("X-Forwarded-For"); prior != "" {
|
||||
clientIP = prior + ", " + clientIP
|
||||
}
|
||||
proxyReq.Header.Set("X-Forwarded-For", clientIP)
|
||||
proxyReq.Header.Set("X-Forwarded-Proto", getScheme(originalReq))
|
||||
proxyReq.Header.Set("X-Forwarded-Host", originalReq.Host)
|
||||
|
||||
// Add custom headers from config
|
||||
for key, value := range rp.config.Headers {
|
||||
// Substitute parameters like {version}
|
||||
substituted := value
|
||||
for paramKey, paramValue := range params {
|
||||
substituted = strings.ReplaceAll(substituted, "{"+paramKey+"}", paramValue)
|
||||
}
|
||||
// Substitute $remote_addr
|
||||
substituted = strings.ReplaceAll(substituted, "$remote_addr", clientIP)
|
||||
proxyReq.Header.Set(key, substituted)
|
||||
}
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) copyResponse(w http.ResponseWriter, resp *http.Response) {
|
||||
// Copy headers
|
||||
for key, values := range resp.Header {
|
||||
for _, value := range values {
|
||||
w.Header().Add(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// Remove hop-by-hop headers from response
|
||||
removeHopByHopHeaders(w.Header())
|
||||
|
||||
// Write status code
|
||||
w.WriteHeader(resp.StatusCode)
|
||||
|
||||
// Copy body
|
||||
io.Copy(w, resp.Body)
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) handleError(w http.ResponseWriter, status int, message string, err error) {
|
||||
if rp.logger != nil {
|
||||
rp.logger.Error(message, "error", err)
|
||||
}
|
||||
http.Error(w, message, status)
|
||||
}
|
||||
|
||||
func (rp *ReverseProxy) handleProxyError(w http.ResponseWriter, err error) {
|
||||
if rp.logger != nil {
|
||||
rp.logger.Error("Proxy request failed", "error", err)
|
||||
}
|
||||
|
||||
// Check for timeout
|
||||
if err, ok := err.(net.Error); ok && err.Timeout() {
|
||||
http.Error(w, "504 Gateway Timeout", http.StatusGatewayTimeout)
|
||||
return
|
||||
}
|
||||
|
||||
// Check for connection errors
|
||||
if isConnectionError(err) {
|
||||
http.Error(w, "502 Bad Gateway", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
|
||||
// Context cancelled (client disconnected)
|
||||
if err == context.Canceled {
|
||||
return
|
||||
}
|
||||
|
||||
http.Error(w, "502 Bad Gateway", http.StatusBadGateway)
|
||||
}
|
||||
|
||||
// Helper functions
|
||||
|
||||
func singleJoiningSlash(a, b string) string {
|
||||
aslash := strings.HasSuffix(a, "/")
|
||||
bslash := strings.HasPrefix(b, "/")
|
||||
switch {
|
||||
case aslash && bslash:
|
||||
return a + b[1:]
|
||||
case !aslash && !bslash:
|
||||
return a + "/" + b
|
||||
}
|
||||
return a + b
|
||||
}
|
||||
|
||||
func removeHopByHopHeaders(h http.Header) {
|
||||
hopByHopHeaders := []string{
|
||||
"Connection",
|
||||
"Proxy-Connection",
|
||||
"Keep-Alive",
|
||||
"Proxy-Authenticate",
|
||||
"Proxy-Authorization",
|
||||
"Te",
|
||||
"Trailer",
|
||||
"Transfer-Encoding",
|
||||
"Upgrade",
|
||||
}
|
||||
|
||||
for _, header := range hopByHopHeaders {
|
||||
h.Del(header)
|
||||
}
|
||||
}
|
||||
|
||||
func getClientIP(r *http.Request) string {
|
||||
// Check X-Real-IP first
|
||||
if ip := r.Header.Get("X-Real-IP"); ip != "" {
|
||||
return ip
|
||||
}
|
||||
|
||||
// Get from RemoteAddr
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
return r.RemoteAddr
|
||||
}
|
||||
return host
|
||||
}
|
||||
|
||||
func getScheme(r *http.Request) string {
|
||||
if r.TLS != nil {
|
||||
return "https"
|
||||
}
|
||||
if scheme := r.Header.Get("X-Forwarded-Proto"); scheme != "" {
|
||||
return scheme
|
||||
}
|
||||
return "http"
|
||||
}
|
||||
|
||||
func isConnectionError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
errStr := err.Error()
|
||||
connectionErrors := []string{
|
||||
"connection refused",
|
||||
"no such host",
|
||||
"network is unreachable",
|
||||
"connection reset",
|
||||
"broken pipe",
|
||||
}
|
||||
|
||||
for _, connErr := range connectionErrors {
|
||||
if strings.Contains(strings.ToLower(errStr), connErr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,747 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ============== Test Backend Server ==============
|
||||
|
||||
type testBackend struct {
|
||||
server *httptest.Server
|
||||
requestLog []requestLogEntry
|
||||
mu sync.Mutex
|
||||
requestCount int64
|
||||
}
|
||||
|
||||
type requestLogEntry struct {
|
||||
Method string
|
||||
Path string
|
||||
Query string
|
||||
Headers http.Header
|
||||
Body string
|
||||
}
|
||||
|
||||
func newTestBackend(handler http.HandlerFunc) *testBackend {
|
||||
tb := &testBackend{
|
||||
requestLog: make([]requestLogEntry, 0),
|
||||
}
|
||||
|
||||
if handler == nil {
|
||||
handler = tb.defaultHandler
|
||||
}
|
||||
|
||||
tb.server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
tb.logRequest(r)
|
||||
handler(w, r)
|
||||
}))
|
||||
|
||||
return tb
|
||||
}
|
||||
|
||||
func (tb *testBackend) logRequest(r *http.Request) {
|
||||
tb.mu.Lock()
|
||||
defer tb.mu.Unlock()
|
||||
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
// Restore the body for the handler
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
|
||||
tb.requestLog = append(tb.requestLog, requestLogEntry{
|
||||
Method: r.Method,
|
||||
Path: r.URL.Path,
|
||||
Query: r.URL.RawQuery,
|
||||
Headers: r.Header.Clone(),
|
||||
Body: string(body),
|
||||
})
|
||||
atomic.AddInt64(&tb.requestCount, 1)
|
||||
}
|
||||
|
||||
func (tb *testBackend) defaultHandler(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"message": "Backend response",
|
||||
"path": r.URL.Path,
|
||||
"method": r.Method,
|
||||
})
|
||||
}
|
||||
|
||||
func (tb *testBackend) close() {
|
||||
tb.server.Close()
|
||||
}
|
||||
|
||||
func (tb *testBackend) URL() string {
|
||||
return tb.server.URL
|
||||
}
|
||||
|
||||
func (tb *testBackend) getRequestCount() int64 {
|
||||
return atomic.LoadInt64(&tb.requestCount)
|
||||
}
|
||||
|
||||
func (tb *testBackend) getLastRequest() *requestLogEntry {
|
||||
tb.mu.Lock()
|
||||
defer tb.mu.Unlock()
|
||||
if len(tb.requestLog) == 0 {
|
||||
return nil
|
||||
}
|
||||
return &tb.requestLog[len(tb.requestLog)-1]
|
||||
}
|
||||
|
||||
// ============== Proxy Creation Tests ==============
|
||||
|
||||
func TestNew_ValidConfig(t *testing.T) {
|
||||
cfg := &Config{
|
||||
Target: "http://localhost:8080",
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
proxy, err := New(cfg, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create proxy: %v", err)
|
||||
}
|
||||
|
||||
if proxy == nil {
|
||||
t.Fatal("Expected proxy instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_EmptyTarget(t *testing.T) {
|
||||
cfg := &Config{
|
||||
Target: "",
|
||||
}
|
||||
|
||||
_, err := New(cfg, nil)
|
||||
if err == nil {
|
||||
t.Error("Expected error for empty target")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_InvalidTargetURL(t *testing.T) {
|
||||
cfg := &Config{
|
||||
Target: "://invalid-url",
|
||||
}
|
||||
|
||||
_, err := New(cfg, nil)
|
||||
if err == nil {
|
||||
t.Error("Expected error for invalid URL")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_DefaultTimeout(t *testing.T) {
|
||||
cfg := &Config{
|
||||
Target: "http://localhost:8080",
|
||||
}
|
||||
|
||||
proxy, err := New(cfg, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create proxy: %v", err)
|
||||
}
|
||||
|
||||
if proxy.httpClient.Timeout != 30*time.Second {
|
||||
t.Errorf("Expected default timeout 30s, got %v", proxy.httpClient.Timeout)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Basic Proxy Tests ==============
|
||||
|
||||
func TestProxy_BasicGET(t *testing.T) {
|
||||
backend := newTestBackend(nil)
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
|
||||
var response map[string]interface{}
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["path"] != "/test" {
|
||||
t.Errorf("Expected path /test, got %v", response["path"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_BasicPOST(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"received": string(body),
|
||||
"method": r.Method,
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("POST", "/api/data", strings.NewReader(`{"key":"value"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Errorf("Expected status 200, got %d", rr.Code)
|
||||
}
|
||||
|
||||
var response map[string]interface{}
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["method"] != "POST" {
|
||||
t.Errorf("Expected method POST, got %v", response["method"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_PUT(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(map[string]string{"method": r.Method})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("PUT", "/resource/123", strings.NewReader(`{"name":"updated"}`))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["method"] != "PUT" {
|
||||
t.Errorf("Expected method PUT, got %v", response["method"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_DELETE(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(map[string]string{"method": r.Method})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/resource/123", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["method"] != "DELETE" {
|
||||
t.Errorf("Expected method DELETE, got %v", response["method"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Header Tests ==============
|
||||
|
||||
func TestProxy_HeadersForwarding(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"custom_header": r.Header.Get("X-Custom-Header"),
|
||||
"forwarded_for": r.Header.Get("X-Forwarded-For"),
|
||||
"forwarded_host": r.Header.Get("X-Forwarded-Host"),
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/headers", nil)
|
||||
req.Header.Set("X-Custom-Header", "test-value")
|
||||
req.RemoteAddr = "192.168.1.100:12345"
|
||||
req.Host = "example.com"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]interface{}
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["custom_header"] != "test-value" {
|
||||
t.Errorf("Expected custom header, got %v", response["custom_header"])
|
||||
}
|
||||
|
||||
if response["forwarded_for"] != "192.168.1.100" {
|
||||
t.Errorf("Expected X-Forwarded-For, got %v", response["forwarded_for"])
|
||||
}
|
||||
|
||||
if response["forwarded_host"] != "example.com" {
|
||||
t.Errorf("Expected X-Forwarded-Host, got %v", response["forwarded_host"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_CustomHeaders(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"api_version": r.Header.Get("X-API-Version"),
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{
|
||||
Target: backend.URL(),
|
||||
Headers: map[string]string{
|
||||
"X-API-Version": "{version}",
|
||||
},
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/api", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
// Simulate parameter substitution
|
||||
proxy.ProxyRequest(rr, req, map[string]string{"version": "2"})
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["api_version"] != "2" {
|
||||
t.Errorf("Expected API version 2, got %v", response["api_version"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_RemoteAddrSubstitution(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"client_ip": r.Header.Get("X-Client-IP"),
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{
|
||||
Target: backend.URL(),
|
||||
Headers: map[string]string{
|
||||
"X-Client-IP": "$remote_addr",
|
||||
},
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/api", nil)
|
||||
req.RemoteAddr = "10.0.0.1:54321"
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["client_ip"] != "10.0.0.1" {
|
||||
t.Errorf("Expected client IP 10.0.0.1, got %v", response["client_ip"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Query String Tests ==============
|
||||
|
||||
func TestProxy_QueryStringPreservation(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"query": r.URL.RawQuery,
|
||||
"param": r.URL.Query().Get("key"),
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/search?key=value&page=2", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["param"] != "value" {
|
||||
t.Errorf("Expected query param 'value', got %v", response["param"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Status Code Tests ==============
|
||||
|
||||
func TestProxy_StatusCodePreservation(t *testing.T) {
|
||||
statusCodes := []int{200, 201, 400, 404, 500}
|
||||
|
||||
for _, code := range statusCodes {
|
||||
code := code // capture range variable
|
||||
t.Run(fmt.Sprintf("Status_%d", code), func(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(code)
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/status", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != code {
|
||||
t.Errorf("Expected status %d, got %d", code, rr.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Error Handling Tests ==============
|
||||
|
||||
func TestProxy_BackendUnavailable(t *testing.T) {
|
||||
// Use a port that's definitely not listening
|
||||
proxy, _ := New(&Config{
|
||||
Target: "http://127.0.0.1:59999",
|
||||
Timeout: 1 * time.Second,
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusBadGateway {
|
||||
t.Errorf("Expected status 502 Bad Gateway, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_Timeout(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(2 * time.Second)
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{
|
||||
Target: backend.URL(),
|
||||
Timeout: 100 * time.Millisecond,
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/slow", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusGatewayTimeout {
|
||||
t.Errorf("Expected status 504 Gateway Timeout, got %d", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Path Handling Tests ==============
|
||||
|
||||
func TestProxy_StripPrefix(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{
|
||||
Target: backend.URL(),
|
||||
StripPrefix: "/api/v1",
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/api/v1/users", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["path"] != "/users" {
|
||||
t.Errorf("Expected stripped path /users, got %v", response["path"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_TargetWithPath(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{
|
||||
Target: backend.URL() + "/backend",
|
||||
}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/resource", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]string
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["path"] != "/backend/resource" {
|
||||
t.Errorf("Expected path /backend/resource, got %v", response["path"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Large Body Tests ==============
|
||||
|
||||
func TestProxy_LargeRequestBody(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
json.NewEncoder(w).Encode(map[string]int{
|
||||
"received_bytes": len(body),
|
||||
})
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
// 100KB body
|
||||
largeBody := strings.Repeat("x", 100000)
|
||||
req := httptest.NewRequest("POST", "/upload", strings.NewReader(largeBody))
|
||||
req.ContentLength = int64(len(largeBody))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
var response map[string]int
|
||||
json.NewDecoder(rr.Body).Decode(&response)
|
||||
|
||||
if response["received_bytes"] != 100000 {
|
||||
t.Errorf("Expected 100000 bytes, got %d", response["received_bytes"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxy_LargeResponseBody(t *testing.T) {
|
||||
largeResponse := strings.Repeat("y", 100000)
|
||||
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Write([]byte(largeResponse))
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/large", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Body.Len() != 100000 {
|
||||
t.Errorf("Expected 100000 bytes in response, got %d", rr.Body.Len())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Concurrent Requests Tests ==============
|
||||
|
||||
func TestProxy_ConcurrentRequests(t *testing.T) {
|
||||
backend := newTestBackend(nil)
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
const numRequests = 50
|
||||
var wg sync.WaitGroup
|
||||
errors := make(chan error, numRequests)
|
||||
|
||||
for i := 0; i < numRequests; i++ {
|
||||
wg.Add(1)
|
||||
go func(n int) {
|
||||
defer wg.Done()
|
||||
|
||||
req := httptest.NewRequest("GET", "/concurrent", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
errors <- &net.OpError{Op: "test", Err: context.DeadlineExceeded}
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
close(errors)
|
||||
|
||||
errorCount := 0
|
||||
for range errors {
|
||||
errorCount++
|
||||
}
|
||||
|
||||
if errorCount > 0 {
|
||||
t.Errorf("Got %d errors in concurrent requests", errorCount)
|
||||
}
|
||||
|
||||
if backend.getRequestCount() != numRequests {
|
||||
t.Errorf("Expected %d requests at backend, got %d", numRequests, backend.getRequestCount())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Echo Tests ==============
|
||||
|
||||
func TestProxy_Echo(t *testing.T) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
w.Header().Set("Content-Type", r.Header.Get("Content-Type"))
|
||||
w.Write(body)
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
testData := "Hello, Proxy!"
|
||||
req := httptest.NewRequest("POST", "/echo", strings.NewReader(testData))
|
||||
req.Header.Set("Content-Type", "text/plain")
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
proxy.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Body.String() != testData {
|
||||
t.Errorf("Expected echo of '%s', got '%s'", testData, rr.Body.String())
|
||||
}
|
||||
|
||||
if rr.Header().Get("Content-Type") != "text/plain" {
|
||||
t.Errorf("Expected Content-Type text/plain, got %s", rr.Header().Get("Content-Type"))
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Helper Function Tests ==============
|
||||
|
||||
func TestSingleJoiningSlash(t *testing.T) {
|
||||
tests := []struct {
|
||||
a, b, expected string
|
||||
}{
|
||||
{"/api", "/users", "/api/users"},
|
||||
{"/api/", "/users", "/api/users"},
|
||||
{"/api", "users", "/api/users"},
|
||||
{"/api/", "users", "/api/users"},
|
||||
{"", "/users", "/users"},
|
||||
{"/api", "", "/api/"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
result := singleJoiningSlash(tt.a, tt.b)
|
||||
if result != tt.expected {
|
||||
t.Errorf("singleJoiningSlash(%q, %q) = %q, want %q", tt.a, tt.b, result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetClientIP(t *testing.T) {
|
||||
tests := []struct {
|
||||
remoteAddr string
|
||||
xRealIP string
|
||||
expected string
|
||||
}{
|
||||
{"192.168.1.1:1234", "", "192.168.1.1"},
|
||||
{"192.168.1.1:1234", "10.0.0.1", "10.0.0.1"},
|
||||
{"invalid", "", "invalid"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
req.RemoteAddr = tt.remoteAddr
|
||||
if tt.xRealIP != "" {
|
||||
req.Header.Set("X-Real-IP", tt.xRealIP)
|
||||
}
|
||||
|
||||
result := getClientIP(req)
|
||||
if result != tt.expected {
|
||||
t.Errorf("getClientIP() = %q, want %q", result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetScheme(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tls bool
|
||||
header string
|
||||
expected string
|
||||
}{
|
||||
{"HTTP", false, "", "http"},
|
||||
{"HTTPS from TLS", true, "", "https"},
|
||||
{"HTTPS from header", false, "https", "https"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
if tt.header != "" {
|
||||
req.Header.Set("X-Forwarded-Proto", tt.header)
|
||||
}
|
||||
// Note: httptest doesn't set TLS, so we can only test non-TLS cases fully
|
||||
|
||||
result := getScheme(req)
|
||||
if !tt.tls && result != tt.expected {
|
||||
t.Errorf("getScheme() = %q, want %q", result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsConnectionError(t *testing.T) {
|
||||
tests := []struct {
|
||||
err error
|
||||
expected bool
|
||||
}{
|
||||
{nil, false},
|
||||
{&net.OpError{Op: "dial", Err: &net.DNSError{Err: "no such host"}}, true},
|
||||
{context.DeadlineExceeded, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
result := isConnectionError(tt.err)
|
||||
if result != tt.expected {
|
||||
t.Errorf("isConnectionError(%v) = %v, want %v", tt.err, result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Benchmarks ==============
|
||||
|
||||
func BenchmarkProxy_SimpleGET(b *testing.B) {
|
||||
backend := newTestBackend(nil)
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
req := httptest.NewRequest("GET", "/bench", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
proxy.ServeHTTP(rr, req)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkProxy_POSTWithBody(b *testing.B) {
|
||||
backend := newTestBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
io.Copy(io.Discard, r.Body)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
defer backend.close()
|
||||
|
||||
proxy, _ := New(&Config{Target: backend.URL()}, nil)
|
||||
body := strings.Repeat("x", 1024) // 1KB body
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
req := httptest.NewRequest("POST", "/bench", strings.NewReader(body))
|
||||
rr := httptest.NewRecorder()
|
||||
proxy.ServeHTTP(rr, req)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,492 @@
|
||||
// Package routing provides HTTP routing with regex support
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/config"
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
"github.com/konduktor/konduktor/internal/proxy"
|
||||
)
|
||||
|
||||
// RouteMatch represents a matched route with captured parameters
|
||||
type RouteMatch struct {
|
||||
Config map[string]interface{}
|
||||
Params map[string]string
|
||||
}
|
||||
|
||||
// RegexRoute represents a compiled regex route
|
||||
type RegexRoute struct {
|
||||
Pattern *regexp.Regexp
|
||||
Config map[string]interface{}
|
||||
CaseSensitive bool
|
||||
OriginalExpr string
|
||||
}
|
||||
|
||||
// Router handles HTTP routing with exact, regex, and default routes
|
||||
type Router struct {
|
||||
config *config.Config
|
||||
logger *logging.Logger
|
||||
mux *http.ServeMux
|
||||
staticDir string
|
||||
exactRoutes map[string]map[string]interface{}
|
||||
regexRoutes []*RegexRoute
|
||||
defaultRoute map[string]interface{}
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// New creates a new router from config
|
||||
func New(cfg *config.Config, logger *logging.Logger) *Router {
|
||||
staticDir := "./static"
|
||||
if cfg != nil && cfg.HTTP.StaticDir != "" {
|
||||
staticDir = cfg.HTTP.StaticDir
|
||||
}
|
||||
|
||||
r := &Router{
|
||||
config: cfg,
|
||||
logger: logger,
|
||||
mux: http.NewServeMux(),
|
||||
staticDir: staticDir,
|
||||
exactRoutes: make(map[string]map[string]interface{}),
|
||||
regexRoutes: make([]*RegexRoute, 0),
|
||||
}
|
||||
|
||||
// Load routes from extensions
|
||||
if cfg != nil {
|
||||
for _, ext := range cfg.Extensions {
|
||||
if ext.Type == "routing" && ext.Config != nil {
|
||||
if locations, ok := ext.Config["regex_locations"].(map[string]interface{}); ok {
|
||||
for pattern, routeCfg := range locations {
|
||||
if rc, ok := routeCfg.(map[string]interface{}); ok {
|
||||
r.AddRoute(pattern, rc)
|
||||
if logger != nil {
|
||||
logger.Debug("Added route", "pattern", pattern)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
r.setupRoutes()
|
||||
return r
|
||||
}
|
||||
|
||||
// NewRouter creates a router without config (for testing)
|
||||
func NewRouter(opts ...RouterOption) *Router {
|
||||
r := &Router{
|
||||
mux: http.NewServeMux(),
|
||||
staticDir: "./static",
|
||||
exactRoutes: make(map[string]map[string]interface{}),
|
||||
regexRoutes: make([]*RegexRoute, 0),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(r)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// RouterOption is a functional option for Router
|
||||
type RouterOption func(*Router)
|
||||
|
||||
// WithStaticDir sets the static directory
|
||||
func WithStaticDir(dir string) RouterOption {
|
||||
return func(r *Router) {
|
||||
r.staticDir = dir
|
||||
}
|
||||
}
|
||||
|
||||
// StaticDir returns the static directory path
|
||||
func (r *Router) StaticDir() string {
|
||||
return r.staticDir
|
||||
}
|
||||
|
||||
// Routes returns the regex routes (for testing)
|
||||
func (r *Router) Routes() []*RegexRoute {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.regexRoutes
|
||||
}
|
||||
|
||||
// ExactRoutes returns the exact routes (for testing)
|
||||
func (r *Router) ExactRoutes() map[string]map[string]interface{} {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.exactRoutes
|
||||
}
|
||||
|
||||
// DefaultRoute returns the default route (for testing)
|
||||
func (r *Router) DefaultRoute() map[string]interface{} {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.defaultRoute
|
||||
}
|
||||
|
||||
// AddRoute adds a route with the given pattern and config
|
||||
// Pattern formats:
|
||||
// - "=/path" - exact match
|
||||
// - "~regex" - case-sensitive regex
|
||||
// - "~*regex" - case-insensitive regex
|
||||
// - "__default__" - default/fallback route
|
||||
func (r *Router) AddRoute(pattern string, routeConfig map[string]interface{}) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
switch {
|
||||
case pattern == "__default__":
|
||||
r.defaultRoute = routeConfig
|
||||
|
||||
case strings.HasPrefix(pattern, "="):
|
||||
// Exact match route
|
||||
path := strings.TrimPrefix(pattern, "=")
|
||||
r.exactRoutes[path] = routeConfig
|
||||
|
||||
case strings.HasPrefix(pattern, "~*"):
|
||||
// Case-insensitive regex
|
||||
expr := strings.TrimPrefix(pattern, "~*")
|
||||
re, err := regexp.Compile("(?i)" + expr)
|
||||
if err != nil {
|
||||
if r.logger != nil {
|
||||
r.logger.Error("Invalid regex pattern", "pattern", pattern, "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
r.regexRoutes = append(r.regexRoutes, &RegexRoute{
|
||||
Pattern: re,
|
||||
Config: routeConfig,
|
||||
CaseSensitive: false,
|
||||
OriginalExpr: expr,
|
||||
})
|
||||
|
||||
case strings.HasPrefix(pattern, "~"):
|
||||
// Case-sensitive regex
|
||||
expr := strings.TrimPrefix(pattern, "~")
|
||||
re, err := regexp.Compile(expr)
|
||||
if err != nil {
|
||||
if r.logger != nil {
|
||||
r.logger.Error("Invalid regex pattern", "pattern", pattern, "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
r.regexRoutes = append(r.regexRoutes, &RegexRoute{
|
||||
Pattern: re,
|
||||
Config: routeConfig,
|
||||
CaseSensitive: true,
|
||||
OriginalExpr: expr,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Match finds the best matching route for a path
|
||||
// Priority: exact match > regex match > default
|
||||
func (r *Router) Match(path string) *RouteMatch {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
// 1. Check exact routes
|
||||
if cfg, ok := r.exactRoutes[path]; ok {
|
||||
return &RouteMatch{
|
||||
Config: cfg,
|
||||
Params: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Check regex routes
|
||||
for _, route := range r.regexRoutes {
|
||||
match := route.Pattern.FindStringSubmatch(path)
|
||||
if match != nil {
|
||||
params := make(map[string]string)
|
||||
|
||||
// Extract named groups
|
||||
names := route.Pattern.SubexpNames()
|
||||
for i, name := range names {
|
||||
if i > 0 && name != "" && i < len(match) {
|
||||
params[name] = match[i]
|
||||
}
|
||||
}
|
||||
|
||||
return &RouteMatch{
|
||||
Config: route.Config,
|
||||
Params: params,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Check default route
|
||||
if r.defaultRoute != nil {
|
||||
return &RouteMatch{
|
||||
Config: r.defaultRoute,
|
||||
Params: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// setupRoutes configures the routes from config
|
||||
func (r *Router) setupRoutes() {
|
||||
// Health check endpoint
|
||||
r.mux.HandleFunc("/health", r.healthHandler)
|
||||
|
||||
// Setup redirect instructions from config
|
||||
if r.config != nil {
|
||||
for from, to := range r.config.Server.RedirectInstructions {
|
||||
fromPath := from
|
||||
toPath := to
|
||||
r.mux.HandleFunc(fromPath, func(w http.ResponseWriter, req *http.Request) {
|
||||
http.Redirect(w, req, toPath, http.StatusMovedPermanently)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Default handler for all other routes
|
||||
r.mux.HandleFunc("/", r.defaultHandler)
|
||||
}
|
||||
|
||||
// ServeHTTP implements http.Handler
|
||||
func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
r.mux.ServeHTTP(w, req)
|
||||
}
|
||||
|
||||
// healthHandler handles health check requests
|
||||
func (r *Router) healthHandler(w http.ResponseWriter, req *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte("OK"))
|
||||
}
|
||||
|
||||
// defaultHandler handles requests that don't match other routes
|
||||
func (r *Router) defaultHandler(w http.ResponseWriter, req *http.Request) {
|
||||
path := req.URL.Path
|
||||
|
||||
// Try to match against configured routes
|
||||
match := r.Match(path)
|
||||
fmt.Printf("DEBUG defaultHandler: path=%q match=%v defaultRoute=%v\n", path, match != nil, r.defaultRoute != nil)
|
||||
if match != nil {
|
||||
fmt.Printf("DEBUG: matched config: %v\n", match.Config)
|
||||
r.handleRouteMatch(w, req, match)
|
||||
return
|
||||
}
|
||||
|
||||
// Try to serve static file
|
||||
if r.staticDir != "" {
|
||||
// Get absolute path for static dir
|
||||
absStaticDir, err := filepath.Abs(r.staticDir)
|
||||
if err != nil {
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
filePath := filepath.Join(absStaticDir, filepath.Clean("/"+path))
|
||||
cleanPath := filepath.Clean(filePath)
|
||||
|
||||
// Prevent directory traversal - ensure path is within static dir
|
||||
if !strings.HasPrefix(cleanPath+string(filepath.Separator), absStaticDir+string(filepath.Separator)) {
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
// Check if file exists
|
||||
info, err := os.Stat(filePath)
|
||||
if err == nil {
|
||||
if info.IsDir() {
|
||||
// Try index.html
|
||||
indexPath := filepath.Join(filePath, "index.html")
|
||||
if _, err := os.Stat(indexPath); err == nil {
|
||||
http.ServeFile(w, req, indexPath)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
http.ServeFile(w, req, filePath)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 404 Not Found
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
|
||||
// handleRouteMatch handles a matched route
|
||||
func (r *Router) handleRouteMatch(w http.ResponseWriter, req *http.Request, match *RouteMatch) {
|
||||
cfg := match.Config
|
||||
|
||||
// Handle proxy_pass directive
|
||||
if proxyTarget, ok := cfg["proxy_pass"].(string); ok {
|
||||
r.handleProxyPass(w, req, proxyTarget, cfg, match.Params)
|
||||
return
|
||||
}
|
||||
|
||||
// Handle "return" directive
|
||||
if ret, ok := cfg["return"].(string); ok {
|
||||
parts := strings.SplitN(ret, " ", 2)
|
||||
statusCode := 200
|
||||
body := "OK"
|
||||
if len(parts) >= 1 {
|
||||
switch parts[0] {
|
||||
case "200":
|
||||
statusCode = 200
|
||||
case "201":
|
||||
statusCode = 201
|
||||
case "301":
|
||||
statusCode = 301
|
||||
case "302":
|
||||
statusCode = 302
|
||||
case "400":
|
||||
statusCode = 400
|
||||
case "404":
|
||||
statusCode = 404
|
||||
case "500":
|
||||
statusCode = 500
|
||||
}
|
||||
}
|
||||
if len(parts) >= 2 {
|
||||
body = parts[1]
|
||||
}
|
||||
|
||||
if ct, ok := cfg["content_type"].(string); ok {
|
||||
w.Header().Set("Content-Type", ct)
|
||||
} else {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
}
|
||||
|
||||
w.WriteHeader(statusCode)
|
||||
w.Write([]byte(body))
|
||||
return
|
||||
}
|
||||
|
||||
// Handle static files with root
|
||||
if root, ok := cfg["root"].(string); ok {
|
||||
path := req.URL.Path
|
||||
|
||||
if indexFile, ok := cfg["index_file"].(string); ok {
|
||||
if path == "/" || strings.HasSuffix(path, "/") {
|
||||
path = "/" + indexFile
|
||||
}
|
||||
}
|
||||
|
||||
// Get absolute path for root dir
|
||||
absRoot, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
filePath := filepath.Join(absRoot, filepath.Clean("/"+path))
|
||||
cleanPath := filepath.Clean(filePath)
|
||||
|
||||
// DEBUG
|
||||
fmt.Printf("DEBUG: path=%q absRoot=%q filePath=%q cleanPath=%q\n", path, absRoot, filePath, cleanPath)
|
||||
fmt.Printf("DEBUG: check1=%q check2=%q\n", cleanPath+string(filepath.Separator), absRoot+string(filepath.Separator))
|
||||
|
||||
// Prevent directory traversal
|
||||
if !strings.HasPrefix(cleanPath+string(filepath.Separator), absRoot+string(filepath.Separator)) {
|
||||
fmt.Printf("DEBUG: FORBIDDEN - path not within root\n")
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
if cacheControl, ok := cfg["cache_control"].(string); ok {
|
||||
w.Header().Set("Cache-Control", cacheControl)
|
||||
}
|
||||
|
||||
if headers, ok := cfg["headers"].([]interface{}); ok {
|
||||
for _, h := range headers {
|
||||
if header, ok := h.(string); ok {
|
||||
parts := strings.SplitN(header, ": ", 2)
|
||||
if len(parts) == 2 {
|
||||
w.Header().Set(parts[0], parts[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
http.ServeFile(w, req, filePath)
|
||||
return
|
||||
}
|
||||
|
||||
// Handle SPA fallback
|
||||
if spaFallback, ok := cfg["spa_fallback"].(bool); ok && spaFallback {
|
||||
root := r.staticDir
|
||||
if rt, ok := cfg["root"].(string); ok {
|
||||
root = rt
|
||||
}
|
||||
|
||||
indexFile := "index.html"
|
||||
if idx, ok := cfg["index_file"].(string); ok {
|
||||
indexFile = idx
|
||||
}
|
||||
|
||||
filePath := filepath.Join(root, indexFile)
|
||||
http.ServeFile(w, req, filePath)
|
||||
return
|
||||
}
|
||||
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
|
||||
// handleProxyPass proxies the request to the target backend
|
||||
func (r *Router) handleProxyPass(w http.ResponseWriter, req *http.Request, target string, cfg map[string]interface{}, params map[string]string) {
|
||||
// Substitute params in target URL (e.g., {version} -> actual version)
|
||||
for key, value := range params {
|
||||
target = strings.ReplaceAll(target, "{"+key+"}", value)
|
||||
}
|
||||
|
||||
// Create proxy
|
||||
proxyConfig := &proxy.Config{
|
||||
Target: target,
|
||||
Headers: make(map[string]string),
|
||||
}
|
||||
|
||||
// Parse headers from config
|
||||
if headers, ok := cfg["headers"].([]interface{}); ok {
|
||||
for _, h := range headers {
|
||||
if header, ok := h.(string); ok {
|
||||
parts := strings.SplitN(header, ": ", 2)
|
||||
if len(parts) == 2 {
|
||||
// Substitute params in header values
|
||||
headerValue := parts[1]
|
||||
for key, value := range params {
|
||||
headerValue = strings.ReplaceAll(headerValue, "{"+key+"}", value)
|
||||
}
|
||||
proxyConfig.Headers[parts[0]] = headerValue
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
p, err := proxy.New(proxyConfig, r.logger)
|
||||
if err != nil {
|
||||
if r.logger != nil {
|
||||
r.logger.Error("Failed to create proxy", "target", target, "error", err)
|
||||
}
|
||||
http.Error(w, "Bad Gateway", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
|
||||
p.ProxyRequest(w, req, params)
|
||||
}
|
||||
|
||||
// CreateRouterFromConfig creates a router from extension config
|
||||
func CreateRouterFromConfig(cfg map[string]interface{}) *Router {
|
||||
router := NewRouter()
|
||||
|
||||
if locations, ok := cfg["regex_locations"].(map[string]interface{}); ok {
|
||||
for pattern, routeCfg := range locations {
|
||||
if rc, ok := routeCfg.(map[string]interface{}); ok {
|
||||
router.AddRoute(pattern, rc)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return router
|
||||
}
|
||||
@@ -0,0 +1,375 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ============== Router Initialization Tests ==============
|
||||
|
||||
func TestRouter_Initialization(t *testing.T) {
|
||||
router := NewRouter()
|
||||
|
||||
if router.StaticDir() != "./static" {
|
||||
t.Errorf("Expected static dir ./static, got %s", router.StaticDir())
|
||||
}
|
||||
|
||||
if len(router.Routes()) != 0 {
|
||||
t.Error("Expected empty routes")
|
||||
}
|
||||
|
||||
if len(router.ExactRoutes()) != 0 {
|
||||
t.Error("Expected empty exact routes")
|
||||
}
|
||||
|
||||
if router.DefaultRoute() != nil {
|
||||
t.Error("Expected nil default route")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_CustomStaticDir(t *testing.T) {
|
||||
router := NewRouter(WithStaticDir("/custom/path"))
|
||||
|
||||
if router.StaticDir() != "/custom/path" {
|
||||
t.Errorf("Expected static dir /custom/path, got %s", router.StaticDir())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Route Adding Tests ==============
|
||||
|
||||
func TestRouter_AddExactRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"return": "200 OK"}
|
||||
|
||||
router.AddRoute("=/health", config)
|
||||
|
||||
exactRoutes := router.ExactRoutes()
|
||||
if _, ok := exactRoutes["/health"]; !ok {
|
||||
t.Error("Expected /health in exact routes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_AddDefaultRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"spa_fallback": true, "root": "./static"}
|
||||
|
||||
router.AddRoute("__default__", config)
|
||||
|
||||
if router.DefaultRoute() == nil {
|
||||
t.Error("Expected default route to be set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_AddRegexRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"root": "./static"}
|
||||
|
||||
router.AddRoute("~^/api/", config)
|
||||
|
||||
if len(router.Routes()) != 1 {
|
||||
t.Errorf("Expected 1 regex route, got %d", len(router.Routes()))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_AddCaseInsensitiveRegexRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"root": "./static", "cache_control": "public, max-age=3600"}
|
||||
|
||||
router.AddRoute("~*\\.(css|js)$", config)
|
||||
|
||||
if len(router.Routes()) != 1 {
|
||||
t.Errorf("Expected 1 regex route, got %d", len(router.Routes()))
|
||||
}
|
||||
|
||||
if router.Routes()[0].CaseSensitive {
|
||||
t.Error("Expected case-insensitive route")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_InvalidRegexPattern(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"root": "./static"}
|
||||
|
||||
// Invalid regex - unmatched bracket
|
||||
router.AddRoute("~^/api/[invalid", config)
|
||||
|
||||
// Should not add invalid pattern
|
||||
if len(router.Routes()) != 0 {
|
||||
t.Error("Should not add invalid regex pattern")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Route Matching Tests ==============
|
||||
|
||||
func TestRouter_MatchExactRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"return": "200 OK"}
|
||||
router.AddRoute("=/health", config)
|
||||
|
||||
match := router.Match("/health")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected match for /health")
|
||||
}
|
||||
|
||||
if match.Config["return"] != "200 OK" {
|
||||
t.Error("Expected return config")
|
||||
}
|
||||
|
||||
if len(match.Params) != 0 {
|
||||
t.Error("Expected empty params for exact match")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchExactRouteNoMatch(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"return": "200 OK"}
|
||||
router.AddRoute("=/health", config)
|
||||
|
||||
match := router.Match("/healthcheck")
|
||||
|
||||
if match != nil {
|
||||
t.Error("Exact route should not match /healthcheck")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchRegexRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"proxy_pass": "http://localhost:9001"}
|
||||
router.AddRoute("~^/api/v\\d+/", config)
|
||||
|
||||
match := router.Match("/api/v1/users")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected match for /api/v1/users")
|
||||
}
|
||||
|
||||
if match.Config["proxy_pass"] != "http://localhost:9001" {
|
||||
t.Error("Expected proxy_pass config")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchRegexRouteWithGroups(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"proxy_pass": "http://localhost:9001"}
|
||||
router.AddRoute("~^/api/v(?P<version>\\d+)/", config)
|
||||
|
||||
match := router.Match("/api/v2/data")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected match for /api/v2/data")
|
||||
}
|
||||
|
||||
if match.Params["version"] != "2" {
|
||||
t.Errorf("Expected version=2, got %s", match.Params["version"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchCaseInsensitiveRegex(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"root": "./static", "cache_control": "public, max-age=3600"}
|
||||
router.AddRoute("~*\\.(CSS|JS)$", config)
|
||||
|
||||
// Should match lowercase
|
||||
match1 := router.Match("/styles/main.css")
|
||||
if match1 == nil {
|
||||
t.Error("Should match lowercase .css")
|
||||
}
|
||||
|
||||
// Should match uppercase
|
||||
match2 := router.Match("/scripts/app.JS")
|
||||
if match2 == nil {
|
||||
t.Error("Should match uppercase .JS")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchCaseSensitiveRegex(t *testing.T) {
|
||||
router := NewRouter()
|
||||
config := map[string]interface{}{"root": "./static"}
|
||||
router.AddRoute("~\\.(css)$", config)
|
||||
|
||||
// Should match lowercase
|
||||
match1 := router.Match("/styles/main.css")
|
||||
if match1 == nil {
|
||||
t.Error("Should match lowercase .css")
|
||||
}
|
||||
|
||||
// Should NOT match uppercase
|
||||
match2 := router.Match("/styles/main.CSS")
|
||||
if match2 != nil {
|
||||
t.Error("Should not match uppercase .CSS for case-sensitive regex")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_MatchDefaultRoute(t *testing.T) {
|
||||
router := NewRouter()
|
||||
router.AddRoute("=/health", map[string]interface{}{"return": "200 OK"})
|
||||
router.AddRoute("__default__", map[string]interface{}{"spa_fallback": true})
|
||||
|
||||
match := router.Match("/unknown/path")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected default route match")
|
||||
}
|
||||
|
||||
if match.Config["spa_fallback"] != true {
|
||||
t.Error("Expected spa_fallback config from default route")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Priority Tests ==============
|
||||
|
||||
func TestRouter_PriorityExactOverRegex(t *testing.T) {
|
||||
router := NewRouter()
|
||||
router.AddRoute("=/api/status", map[string]interface{}{"return": "200 Exact"})
|
||||
router.AddRoute("~^/api/", map[string]interface{}{"proxy_pass": "http://localhost:9001"})
|
||||
|
||||
match := router.Match("/api/status")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected match")
|
||||
}
|
||||
|
||||
if match.Config["return"] != "200 Exact" {
|
||||
t.Error("Exact match should have priority over regex")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouter_PriorityRegexOverDefault(t *testing.T) {
|
||||
router := NewRouter()
|
||||
router.AddRoute("~^/api/", map[string]interface{}{"proxy_pass": "http://localhost:9001"})
|
||||
router.AddRoute("__default__", map[string]interface{}{"spa_fallback": true})
|
||||
|
||||
match := router.Match("/api/v1/users")
|
||||
|
||||
if match == nil {
|
||||
t.Fatal("Expected match")
|
||||
}
|
||||
|
||||
if match.Config["proxy_pass"] != "http://localhost:9001" {
|
||||
t.Error("Regex match should have priority over default")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== CreateRouterFromConfig Tests ==============
|
||||
|
||||
func TestCreateRouterFromConfig(t *testing.T) {
|
||||
config := map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"=/health": map[string]interface{}{
|
||||
"return": "200 OK",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"~^/api/": map[string]interface{}{
|
||||
"proxy_pass": "http://localhost:9001",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"spa_fallback": true,
|
||||
"root": "./static",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
router := CreateRouterFromConfig(config)
|
||||
|
||||
// Check exact route
|
||||
if _, ok := router.ExactRoutes()["/health"]; !ok {
|
||||
t.Error("Expected /health exact route")
|
||||
}
|
||||
|
||||
// Check regex route
|
||||
if len(router.Routes()) != 1 {
|
||||
t.Errorf("Expected 1 regex route, got %d", len(router.Routes()))
|
||||
}
|
||||
|
||||
// Check default route
|
||||
if router.DefaultRoute() == nil {
|
||||
t.Error("Expected default route")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Static Dir Path Tests ==============
|
||||
|
||||
func TestRouter_StaticDirPath(t *testing.T) {
|
||||
router := NewRouter(WithStaticDir("/var/www/html"))
|
||||
|
||||
expected, _ := filepath.Abs("/var/www/html")
|
||||
actual, _ := filepath.Abs(router.StaticDir())
|
||||
|
||||
if actual != expected {
|
||||
t.Errorf("Expected static dir %s, got %s", expected, actual)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Concurrent Access Tests ==============
|
||||
|
||||
func TestRouter_ConcurrentAccess(t *testing.T) {
|
||||
router := NewRouter()
|
||||
|
||||
// Add routes concurrently
|
||||
done := make(chan bool, 10)
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
go func(n int) {
|
||||
router.AddRoute("~^/api/v"+string(rune('0'+n))+"/", map[string]interface{}{
|
||||
"proxy_pass": "http://localhost:900" + string(rune('0'+n)),
|
||||
})
|
||||
done <- true
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Wait for all goroutines
|
||||
for i := 0; i < 10; i++ {
|
||||
<-done
|
||||
}
|
||||
|
||||
// Match routes concurrently
|
||||
for i := 0; i < 10; i++ {
|
||||
go func(n int) {
|
||||
router.Match("/api/v" + string(rune('0'+n)) + "/users")
|
||||
done <- true
|
||||
}(i)
|
||||
}
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
<-done
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Benchmarks ==============
|
||||
|
||||
func BenchmarkRouter_MatchExact(b *testing.B) {
|
||||
router := NewRouter()
|
||||
router.AddRoute("=/health", map[string]interface{}{"return": "200 OK"})
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
router.Match("/health")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkRouter_MatchRegex(b *testing.B) {
|
||||
router := NewRouter()
|
||||
router.AddRoute("~^/api/v(?P<version>\\d+)/", map[string]interface{}{"proxy_pass": "http://localhost:9001"})
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
router.Match("/api/v1/users/123")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkRouter_MatchWithManyRoutes(b *testing.B) {
|
||||
router := NewRouter()
|
||||
|
||||
// Add many routes
|
||||
for i := 0; i < 50; i++ {
|
||||
router.AddRoute("~^/api/v"+string(rune('0'+i%10))+"/service"+string(rune('0'+i/10))+"/",
|
||||
map[string]interface{}{"proxy_pass": "http://localhost:9001"})
|
||||
}
|
||||
router.AddRoute("__default__", map[string]interface{}{"spa_fallback": true})
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
router.Match("/api/v5/service3/users/123")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Package server provides the HTTP server implementation
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/config"
|
||||
"github.com/konduktor/konduktor/internal/extension"
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
"github.com/konduktor/konduktor/internal/middleware"
|
||||
)
|
||||
|
||||
const Version = "0.2.0"
|
||||
|
||||
// Server represents the Konduktor HTTP server
|
||||
type Server struct {
|
||||
config *config.Config
|
||||
httpServer *http.Server
|
||||
extensionManager *extension.Manager
|
||||
logger *logging.Logger
|
||||
}
|
||||
|
||||
// New creates a new server instance
|
||||
func New(cfg *config.Config) (*Server, error) {
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return nil, fmt.Errorf("invalid configuration: %w", err)
|
||||
}
|
||||
|
||||
logger, err := logging.NewFromConfig(cfg.Logging)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create logger: %w", err)
|
||||
}
|
||||
|
||||
// Create extension manager
|
||||
extManager := extension.NewManager(logger)
|
||||
|
||||
// Load extensions from config
|
||||
for _, extCfg := range cfg.Extensions {
|
||||
// Add static_dir to routing config if not present
|
||||
if extCfg.Type == "routing" {
|
||||
if extCfg.Config == nil {
|
||||
extCfg.Config = make(map[string]interface{})
|
||||
}
|
||||
if _, ok := extCfg.Config["static_dir"]; !ok {
|
||||
extCfg.Config["static_dir"] = cfg.HTTP.StaticDir
|
||||
}
|
||||
}
|
||||
|
||||
if err := extManager.LoadExtension(extCfg.Type, extCfg.Config); err != nil {
|
||||
logger.Error("Failed to load extension", "type", extCfg.Type, "error", err)
|
||||
// Continue loading other extensions
|
||||
}
|
||||
}
|
||||
|
||||
srv := &Server{
|
||||
config: cfg,
|
||||
extensionManager: extManager,
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
return srv, nil
|
||||
}
|
||||
|
||||
// Run starts the server and blocks until shutdown
|
||||
func (s *Server) Run() error {
|
||||
// Build handler chain with middleware
|
||||
handler := s.buildHandler()
|
||||
|
||||
// Create HTTP server
|
||||
addr := fmt.Sprintf("%s:%d", s.config.Server.Host, s.config.Server.Port)
|
||||
s.httpServer = &http.Server{
|
||||
Addr: addr,
|
||||
Handler: handler,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
WriteTimeout: 30 * time.Second,
|
||||
IdleTimeout: 120 * time.Second,
|
||||
}
|
||||
|
||||
// Start server in goroutine
|
||||
errChan := make(chan error, 1)
|
||||
go func() {
|
||||
s.logger.Info("Server starting", "addr", addr, "version", Version)
|
||||
|
||||
var err error
|
||||
if s.config.SSL.Enabled {
|
||||
err = s.httpServer.ListenAndServeTLS(s.config.SSL.CertFile, s.config.SSL.KeyFile)
|
||||
} else {
|
||||
err = s.httpServer.ListenAndServe()
|
||||
}
|
||||
|
||||
if err != nil && err != http.ErrServerClosed {
|
||||
errChan <- err
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for shutdown signal
|
||||
return s.waitForShutdown(errChan)
|
||||
}
|
||||
|
||||
// buildHandler builds the HTTP handler chain
|
||||
func (s *Server) buildHandler() http.Handler {
|
||||
// Create base handler that returns 404
|
||||
baseHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.NotFound(w, r)
|
||||
})
|
||||
|
||||
// Wrap with extension manager
|
||||
var handler http.Handler = s.extensionManager.Handler(baseHandler)
|
||||
|
||||
// Add middleware (applied in reverse order)
|
||||
handler = middleware.AccessLog(handler, s.logger)
|
||||
handler = middleware.ServerHeader(handler, Version)
|
||||
handler = middleware.Recovery(handler, s.logger)
|
||||
|
||||
return handler
|
||||
}
|
||||
|
||||
// waitForShutdown waits for shutdown signal and gracefully stops the server
|
||||
func (s *Server) waitForShutdown(errChan <-chan error) error {
|
||||
// Listen for shutdown signals
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM)
|
||||
|
||||
select {
|
||||
case err := <-errChan:
|
||||
return err
|
||||
case sig := <-sigChan:
|
||||
s.logger.Info("Shutdown signal received", "signal", sig.String())
|
||||
}
|
||||
|
||||
// Graceful shutdown with timeout
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
s.logger.Info("Shutting down server...")
|
||||
|
||||
// Cleanup extensions
|
||||
s.extensionManager.Cleanup()
|
||||
|
||||
if err := s.httpServer.Shutdown(ctx); err != nil {
|
||||
s.logger.Error("Error during shutdown", "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
s.logger.Info("Server stopped gracefully")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Shutdown gracefully shuts down the server
|
||||
func (s *Server) Shutdown(ctx context.Context) error {
|
||||
return s.httpServer.Shutdown(ctx)
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
# Integration Tests
|
||||
|
||||
Интеграционные тесты для Konduktor — полноценное тестирование сервера с реальными HTTP запросами.
|
||||
|
||||
## Отличие от unit-тестов
|
||||
|
||||
| Аспект | Unit-тесты | Интеграционные тесты |
|
||||
|--------|------------|---------------------|
|
||||
| Scope | Отдельный модуль в изоляции | Весь сервер целиком |
|
||||
| Backend | Mock (httptest.Server) | Реальные HTTP серверы |
|
||||
| Config | Программный | YAML конфигурация |
|
||||
| Extensions | Не тестируются | Полная цепочка обработки |
|
||||
|
||||
## Структура тестов
|
||||
|
||||
```
|
||||
tests/integration/
|
||||
├── README.md # Эта документация
|
||||
├── helpers_test.go # Общие хелперы и утилиты
|
||||
├── reverse_proxy_test.go # Тесты reverse proxy
|
||||
├── routing_test.go # Тесты маршрутизации (TODO)
|
||||
├── security_test.go # Тесты security extension (TODO)
|
||||
├── caching_test.go # Тесты caching extension (TODO)
|
||||
└── static_files_test.go # Тесты статических файлов (TODO)
|
||||
```
|
||||
|
||||
## Что тестируют интеграционные тесты
|
||||
|
||||
### 1. Reverse Proxy (`reverse_proxy_test.go`)
|
||||
|
||||
- [x] Базовое проксирование GET/POST/PUT/DELETE
|
||||
- [x] Exact match routes (`=/api/version`)
|
||||
- [x] Regex routes с параметрами (`~^/api/resource/(?P<id>\d+)$`)
|
||||
- [x] Подстановка параметров в target URL (`{id}`, `{tag}`)
|
||||
- [x] Подстановка переменных в заголовки (`$remote_addr`)
|
||||
- [x] Передача заголовков X-Forwarded-For, X-Real-IP
|
||||
- [x] Сохранение query string
|
||||
- [x] Обработка ошибок backend (502, 504)
|
||||
- [x] Таймауты соединения
|
||||
|
||||
### 2. Routing Extension (`routing_test.go`)
|
||||
|
||||
- [x] Приоритет маршрутов (exact > regex > default)
|
||||
- [x] Case-sensitive regex (`~`)
|
||||
- [x] Case-insensitive regex (`~*`)
|
||||
- [x] Default route (`__default__`)
|
||||
- [x] Return directive (`return 200 "OK"`)
|
||||
- [x] Regex с именованными группами
|
||||
- [x] Множественные regex маршруты
|
||||
- [x] Кастомные заголовки в маршрутах
|
||||
- [x] Обработка отсутствия маршрута
|
||||
|
||||
### 3. Security Extension (`security_test.go`)
|
||||
|
||||
- [ ] IP whitelist
|
||||
- [ ] IP blacklist
|
||||
- [ ] CIDR нотация (10.0.0.0/8)
|
||||
- [ ] Security headers (X-Frame-Options, X-Content-Type-Options)
|
||||
- [ ] Rate limiting
|
||||
- [ ] Комбинация с другими extensions
|
||||
|
||||
### 4. Caching Extension (`caching_test.go`)
|
||||
|
||||
- [x] Cache hit/miss
|
||||
- [x] TTL expiration
|
||||
- [x] Pattern-based caching
|
||||
- [x] Cache-Control headers (X-Cache header)
|
||||
- [x] Кэширование только GET запросов
|
||||
- [x] Разные пути = разные ключи кэша
|
||||
- [x] Query string влияет на ключ кэша
|
||||
- [x] Ошибки не кэшируются
|
||||
- [x] Конкурентный доступ к кэшу
|
||||
- [x] Множественные паттерны кэширования
|
||||
|
||||
### 5. Static Files (`static_files_test.go`)
|
||||
|
||||
- [ ] Serving статических файлов
|
||||
- [ ] Index file (index.html)
|
||||
- [ ] MIME types
|
||||
- [ ] Cache-Control для static
|
||||
- [ ] SPA fallback
|
||||
- [ ] Directory traversal protection
|
||||
- [ ] 404 для несуществующих файлов
|
||||
|
||||
### 6. Extension Chain (`extension_chain_test.go`)
|
||||
|
||||
- [ ] Порядок выполнения extensions (security → caching → routing)
|
||||
- [ ] Прерывание цепочки при ошибке
|
||||
- [ ] Совместная работа extensions
|
||||
|
||||
## Запуск тестов
|
||||
|
||||
```bash
|
||||
# Все интеграционные тесты
|
||||
go test ./tests/integration/... -v
|
||||
|
||||
# Конкретный файл
|
||||
go test ./tests/integration/... -v -run TestReverseProxy
|
||||
|
||||
# С таймаутом (интеграционные тесты медленнее)
|
||||
go test ./tests/integration/... -v -timeout 60s
|
||||
|
||||
# С покрытием
|
||||
go test ./tests/integration/... -v -coverprofile=coverage.out
|
||||
```
|
||||
|
||||
## Требования
|
||||
|
||||
- Свободные порты: тесты используют случайные порты (`:0`)
|
||||
- Сетевой доступ: для localhost соединений
|
||||
- Время: интеграционные тесты занимают больше времени (~5-10 сек)
|
||||
|
||||
## Добавление новых тестов
|
||||
|
||||
1. Создайте файл `*_test.go` в `tests/integration/`
|
||||
2. Используйте хелперы из `helpers_test.go`:
|
||||
- `startTestServer()` — запуск Konduktor сервера
|
||||
- `startBackend()` — запуск mock backend
|
||||
- `makeRequest()` — отправка HTTP запроса
|
||||
3. Добавьте описание в этот README
|
||||
|
||||
## CI/CD
|
||||
|
||||
Интеграционные тесты запускаются отдельно от unit-тестов:
|
||||
|
||||
```yaml
|
||||
# .github/workflows/test.yml
|
||||
jobs:
|
||||
unit-tests:
|
||||
run: go test ./internal/...
|
||||
|
||||
integration-tests:
|
||||
run: go test ./tests/integration/... -timeout 120s
|
||||
```
|
||||
@@ -0,0 +1,666 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/extension"
|
||||
)
|
||||
|
||||
// ============== Basic Cache Hit/Miss Tests ==============
|
||||
|
||||
func TestCaching_BasicHitMiss(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
count := atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"request_number": count,
|
||||
"timestamp": time.Now().UnixNano(),
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
// Create caching extension
|
||||
cachingExt, err := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "30s",
|
||||
"methods": []interface{}{"GET"},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create caching extension: %v", err)
|
||||
}
|
||||
|
||||
// Create routing extension
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// First request - should be MISS
|
||||
resp1, err := client.Get("/api/data", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request 1 failed: %v", err)
|
||||
}
|
||||
|
||||
cacheHeader1 := resp1.Header.Get("X-Cache")
|
||||
var result1 map[string]interface{}
|
||||
json.NewDecoder(resp1.Body).Decode(&result1)
|
||||
resp1.Body.Close()
|
||||
|
||||
if cacheHeader1 != "MISS" {
|
||||
t.Errorf("Expected X-Cache: MISS for first request, got %q", cacheHeader1)
|
||||
}
|
||||
|
||||
// Second request - should be HIT (same response)
|
||||
resp2, err := client.Get("/api/data", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request 2 failed: %v", err)
|
||||
}
|
||||
|
||||
cacheHeader2 := resp2.Header.Get("X-Cache")
|
||||
var result2 map[string]interface{}
|
||||
json.NewDecoder(resp2.Body).Decode(&result2)
|
||||
resp2.Body.Close()
|
||||
|
||||
if cacheHeader2 != "HIT" {
|
||||
t.Errorf("Expected X-Cache: HIT for second request, got %q", cacheHeader2)
|
||||
}
|
||||
|
||||
// Verify same response (from cache)
|
||||
if result1["request_number"] != result2["request_number"] {
|
||||
t.Errorf("Expected same request_number from cache, got %v and %v",
|
||||
result1["request_number"], result2["request_number"])
|
||||
}
|
||||
|
||||
// Backend should only receive 1 request
|
||||
if atomic.LoadInt64(&requestCount) != 1 {
|
||||
t.Errorf("Expected 1 backend request, got %d", requestCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== TTL Expiration Tests ==============
|
||||
|
||||
func TestCaching_TTLExpiration(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
count := atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"request_number": count,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
// Create caching extension with short TTL
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "100ms", // Very short TTL for testing
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "100ms",
|
||||
"methods": []interface{}{"GET"},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// First request
|
||||
resp1, _ := client.Get("/api/data", nil)
|
||||
var result1 map[string]interface{}
|
||||
json.NewDecoder(resp1.Body).Decode(&result1)
|
||||
resp1.Body.Close()
|
||||
|
||||
// Second request (within TTL) - should be HIT
|
||||
resp2, _ := client.Get("/api/data", nil)
|
||||
cacheHeader2 := resp2.Header.Get("X-Cache")
|
||||
resp2.Body.Close()
|
||||
|
||||
if cacheHeader2 != "HIT" {
|
||||
t.Errorf("Expected X-Cache: HIT before TTL expires, got %q", cacheHeader2)
|
||||
}
|
||||
|
||||
// Wait for TTL to expire
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
|
||||
// Third request (after TTL) - should be MISS
|
||||
resp3, _ := client.Get("/api/data", nil)
|
||||
cacheHeader3 := resp3.Header.Get("X-Cache")
|
||||
var result3 map[string]interface{}
|
||||
json.NewDecoder(resp3.Body).Decode(&result3)
|
||||
resp3.Body.Close()
|
||||
|
||||
if cacheHeader3 != "MISS" {
|
||||
t.Errorf("Expected X-Cache: MISS after TTL expires, got %q", cacheHeader3)
|
||||
}
|
||||
|
||||
// Verify new request was made (different request_number)
|
||||
if result1["request_number"] == result3["request_number"] {
|
||||
t.Error("Expected different request_number after TTL expiration")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Pattern-Based Caching Tests ==============
|
||||
|
||||
func TestCaching_PatternBasedCaching(t *testing.T) {
|
||||
var apiCount, staticCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if r.URL.Path[:5] == "/api/" {
|
||||
atomic.AddInt64(&apiCount, 1)
|
||||
} else {
|
||||
atomic.AddInt64(&staticCount, 1)
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]string{"path": r.URL.Path})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
// Only cache /api/* paths
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
"methods": []interface{}{"GET"},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Multiple requests to /api/ - should be cached
|
||||
for i := 0; i < 3; i++ {
|
||||
resp, _ := client.Get("/api/users", nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// Multiple requests to /static/ - should NOT be cached (not matching pattern)
|
||||
for i := 0; i < 3; i++ {
|
||||
resp, _ := client.Get("/static/file.js", nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// API should have only 1 request (cached)
|
||||
if atomic.LoadInt64(&apiCount) != 1 {
|
||||
t.Errorf("Expected 1 API request (cached), got %d", apiCount)
|
||||
}
|
||||
|
||||
// Static should have 3 requests (not cached)
|
||||
if atomic.LoadInt64(&staticCount) != 3 {
|
||||
t.Errorf("Expected 3 static requests (not cached), got %d", staticCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Method-Specific Caching Tests ==============
|
||||
|
||||
func TestCaching_OnlyGETMethodCached(t *testing.T) {
|
||||
var getCount, postCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if r.Method == "GET" {
|
||||
atomic.AddInt64(&getCount, 1)
|
||||
} else if r.Method == "POST" {
|
||||
atomic.AddInt64(&postCount, 1)
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"method": r.Method,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
"methods": []interface{}{"GET"}, // Only GET
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Multiple GET requests - should be cached
|
||||
for i := 0; i < 3; i++ {
|
||||
resp, _ := client.Get("/api/data", nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// Multiple POST requests - should NOT be cached
|
||||
for i := 0; i < 3; i++ {
|
||||
resp, _ := client.Post("/api/data", []byte(`{}`), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
})
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
if atomic.LoadInt64(&getCount) != 1 {
|
||||
t.Errorf("Expected 1 GET request (cached), got %d", getCount)
|
||||
}
|
||||
|
||||
if atomic.LoadInt64(&postCount) != 3 {
|
||||
t.Errorf("Expected 3 POST requests (not cached), got %d", postCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Different Paths Different Cache Keys ==============
|
||||
|
||||
func TestCaching_DifferentPathsDifferentCacheKeys(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
count := atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"path": r.URL.Path,
|
||||
"request_number": count,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Request different paths
|
||||
paths := []string{"/api/users", "/api/posts", "/api/comments"}
|
||||
|
||||
for _, path := range paths {
|
||||
resp, _ := client.Get(path, nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// Each path should result in a separate backend request
|
||||
if atomic.LoadInt64(&requestCount) != 3 {
|
||||
t.Errorf("Expected 3 backend requests (one per path), got %d", requestCount)
|
||||
}
|
||||
|
||||
// Request same paths again - all should be cached
|
||||
for _, path := range paths {
|
||||
resp, _ := client.Get(path, nil)
|
||||
cacheHeader := resp.Header.Get("X-Cache")
|
||||
resp.Body.Close()
|
||||
|
||||
if cacheHeader != "HIT" {
|
||||
t.Errorf("Expected X-Cache: HIT for %s, got %q", path, cacheHeader)
|
||||
}
|
||||
}
|
||||
|
||||
// No additional backend requests
|
||||
if atomic.LoadInt64(&requestCount) != 3 {
|
||||
t.Errorf("Expected still 3 backend requests after cache hits, got %d", requestCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Query String Affects Cache Key ==============
|
||||
|
||||
func TestCaching_QueryStringAffectsCacheKey(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
count := atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"query": r.URL.RawQuery,
|
||||
"request_number": count,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Different query strings = different cache keys
|
||||
queries := []string{
|
||||
"/api/search?q=hello",
|
||||
"/api/search?q=world",
|
||||
"/api/search?q=test",
|
||||
}
|
||||
|
||||
for _, query := range queries {
|
||||
resp, _ := client.Get(query, nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// Each unique query should result in a separate backend request
|
||||
if atomic.LoadInt64(&requestCount) != 3 {
|
||||
t.Errorf("Expected 3 backend requests (one per query), got %d", requestCount)
|
||||
}
|
||||
|
||||
// Same query again should be cached
|
||||
resp, _ := client.Get("/api/search?q=hello", nil)
|
||||
cacheHeader := resp.Header.Get("X-Cache")
|
||||
resp.Body.Close()
|
||||
|
||||
if cacheHeader != "HIT" {
|
||||
t.Errorf("Expected X-Cache: HIT for repeated query, got %q", cacheHeader)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Cache Does Not Store Error Responses ==============
|
||||
|
||||
func TestCaching_DoesNotCacheErrors(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
json.NewEncoder(w).Encode(map[string]string{"error": "internal error"})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Multiple requests to error endpoint
|
||||
for i := 0; i < 3; i++ {
|
||||
resp, _ := client.Get("/api/error", nil)
|
||||
resp.Body.Close()
|
||||
}
|
||||
|
||||
// All requests should reach backend (errors not cached)
|
||||
if atomic.LoadInt64(&requestCount) != 3 {
|
||||
t.Errorf("Expected 3 backend requests (errors not cached), got %d", requestCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Concurrent Cache Access ==============
|
||||
|
||||
func TestCaching_ConcurrentAccess(t *testing.T) {
|
||||
var requestCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Small delay to increase chance of race conditions
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
count := atomic.AddInt64(&requestCount, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"request_number": count,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "1m",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
const numRequests = 20
|
||||
results := make(chan error, numRequests)
|
||||
|
||||
// Make first request to populate cache
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, _ := client.Get("/api/concurrent", nil)
|
||||
resp.Body.Close()
|
||||
|
||||
// Now many concurrent requests should all hit cache
|
||||
for i := 0; i < numRequests; i++ {
|
||||
go func(n int) {
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/api/concurrent", nil)
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
|
||||
cacheHeader := resp.Header.Get("X-Cache")
|
||||
resp.Body.Close()
|
||||
|
||||
if cacheHeader != "HIT" {
|
||||
results <- fmt.Errorf("request %d: expected HIT, got %s", n, cacheHeader)
|
||||
return
|
||||
}
|
||||
results <- nil
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Collect results
|
||||
var errors []error
|
||||
for i := 0; i < numRequests; i++ {
|
||||
if err := <-results; err != nil {
|
||||
errors = append(errors, err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(errors) > 0 {
|
||||
t.Errorf("Got %d errors in concurrent cache access: %v", len(errors), errors[:min(5, len(errors))])
|
||||
}
|
||||
|
||||
// Only 1 request should reach backend (the initial one)
|
||||
if atomic.LoadInt64(&requestCount) != 1 {
|
||||
t.Errorf("Expected 1 backend request, got %d", requestCount)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Multiple Cache Patterns ==============
|
||||
|
||||
func TestCaching_MultipleCachePatterns(t *testing.T) {
|
||||
var apiCount, staticCount int64
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if len(r.URL.Path) >= 5 && r.URL.Path[:5] == "/api/" {
|
||||
atomic.AddInt64(&apiCount, 1)
|
||||
} else if len(r.URL.Path) >= 8 && r.URL.Path[:8] == "/static/" {
|
||||
atomic.AddInt64(&staticCount, 1)
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]string{"path": r.URL.Path})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
|
||||
cachingExt, _ := extension.NewCachingExtension(map[string]interface{}{
|
||||
"default_ttl": "1m",
|
||||
"cache_patterns": []interface{}{
|
||||
map[string]interface{}{
|
||||
"pattern": "^/api/.*",
|
||||
"ttl": "30s",
|
||||
"methods": []interface{}{"GET"},
|
||||
},
|
||||
map[string]interface{}{
|
||||
"pattern": "^/static/.*",
|
||||
"ttl": "1h", // Static files cached longer
|
||||
"methods": []interface{}{"GET"},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{cachingExt, routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Multiple requests to both patterns
|
||||
for i := 0; i < 3; i++ {
|
||||
resp1, _ := client.Get("/api/data", nil)
|
||||
resp1.Body.Close()
|
||||
|
||||
resp2, _ := client.Get("/static/app.js", nil)
|
||||
resp2.Body.Close()
|
||||
}
|
||||
|
||||
// Both should be cached (1 request each)
|
||||
if atomic.LoadInt64(&apiCount) != 1 {
|
||||
t.Errorf("Expected 1 API request, got %d", apiCount)
|
||||
}
|
||||
|
||||
if atomic.LoadInt64(&staticCount) != 1 {
|
||||
t.Errorf("Expected 1 static request, got %d", staticCount)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/extension"
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
"github.com/konduktor/konduktor/internal/middleware"
|
||||
)
|
||||
|
||||
// TestServer represents a running Konduktor server for testing
|
||||
type TestServer struct {
|
||||
Server *http.Server
|
||||
URL string
|
||||
Port int
|
||||
listener net.Listener
|
||||
handler http.Handler
|
||||
t *testing.T
|
||||
}
|
||||
|
||||
// TestBackend represents a mock backend server
|
||||
type TestBackend struct {
|
||||
server *httptest.Server
|
||||
requestLog []RequestLogEntry
|
||||
mu sync.Mutex
|
||||
requestCount int64
|
||||
handler http.HandlerFunc
|
||||
}
|
||||
|
||||
// RequestLogEntry stores information about a received request
|
||||
type RequestLogEntry struct {
|
||||
Method string
|
||||
Path string
|
||||
Query string
|
||||
Headers http.Header
|
||||
Body string
|
||||
Timestamp time.Time
|
||||
}
|
||||
|
||||
// ServerConfig holds configuration for starting a test server
|
||||
type ServerConfig struct {
|
||||
Extensions []extension.Extension
|
||||
StaticDir string
|
||||
Middleware []func(http.Handler) http.Handler
|
||||
}
|
||||
|
||||
// ============== Test Server ==============
|
||||
|
||||
// StartTestServer creates and starts a Konduktor server for testing
|
||||
func StartTestServer(t *testing.T, cfg *ServerConfig) *TestServer {
|
||||
t.Helper()
|
||||
|
||||
logger, err := logging.New(logging.Config{
|
||||
Level: "DEBUG",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create logger: %v", err)
|
||||
}
|
||||
|
||||
// Create extension manager
|
||||
extManager := extension.NewManager(logger)
|
||||
|
||||
// Add extensions if provided
|
||||
if cfg != nil && len(cfg.Extensions) > 0 {
|
||||
for _, ext := range cfg.Extensions {
|
||||
if err := extManager.AddExtension(ext); err != nil {
|
||||
t.Fatalf("Failed to add extension: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Create a fallback handler for when no extension handles the request
|
||||
fallback := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.NotFound(w, r)
|
||||
})
|
||||
|
||||
// Create handler chain
|
||||
var handler http.Handler = extManager.Handler(fallback)
|
||||
|
||||
// Add middleware
|
||||
handler = middleware.AccessLog(handler, logger)
|
||||
handler = middleware.Recovery(handler, logger)
|
||||
|
||||
// Add custom middleware if provided
|
||||
if cfg != nil {
|
||||
for _, mw := range cfg.Middleware {
|
||||
handler = mw(handler)
|
||||
}
|
||||
}
|
||||
|
||||
// Find available port
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to find available port: %v", err)
|
||||
}
|
||||
|
||||
port := listener.Addr().(*net.TCPAddr).Port
|
||||
|
||||
server := &http.Server{
|
||||
Handler: handler,
|
||||
ReadTimeout: 10 * time.Second,
|
||||
WriteTimeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
ts := &TestServer{
|
||||
Server: server,
|
||||
URL: fmt.Sprintf("http://127.0.0.1:%d", port),
|
||||
Port: port,
|
||||
listener: listener,
|
||||
handler: handler,
|
||||
t: t,
|
||||
}
|
||||
|
||||
// Start server in goroutine
|
||||
go func() {
|
||||
if err := server.Serve(listener); err != nil && err != http.ErrServerClosed {
|
||||
// Don't fail test here as server might be intentionally closed
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for server to be ready
|
||||
ts.waitReady()
|
||||
|
||||
return ts
|
||||
}
|
||||
|
||||
// waitReady waits for the server to be ready to accept connections
|
||||
func (ts *TestServer) waitReady() {
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", ts.Port), 100*time.Millisecond)
|
||||
if err == nil {
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
ts.t.Fatal("Server failed to start within timeout")
|
||||
}
|
||||
|
||||
// Close shuts down the test server
|
||||
func (ts *TestServer) Close() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ts.Server.Shutdown(ctx)
|
||||
}
|
||||
|
||||
// ============== Test Backend ==============
|
||||
|
||||
// StartBackend creates and starts a mock backend server
|
||||
func StartBackend(handler http.HandlerFunc) *TestBackend {
|
||||
tb := &TestBackend{
|
||||
requestLog: make([]RequestLogEntry, 0),
|
||||
handler: handler,
|
||||
}
|
||||
|
||||
if handler == nil {
|
||||
handler = tb.defaultHandler
|
||||
}
|
||||
|
||||
tb.server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
tb.logRequest(r)
|
||||
handler(w, r)
|
||||
}))
|
||||
|
||||
return tb
|
||||
}
|
||||
|
||||
func (tb *TestBackend) logRequest(r *http.Request) {
|
||||
tb.mu.Lock()
|
||||
defer tb.mu.Unlock()
|
||||
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
|
||||
tb.requestLog = append(tb.requestLog, RequestLogEntry{
|
||||
Method: r.Method,
|
||||
Path: r.URL.Path,
|
||||
Query: r.URL.RawQuery,
|
||||
Headers: r.Header.Clone(),
|
||||
Body: string(body),
|
||||
Timestamp: time.Now(),
|
||||
})
|
||||
atomic.AddInt64(&tb.requestCount, 1)
|
||||
}
|
||||
|
||||
func (tb *TestBackend) defaultHandler(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"backend": "default",
|
||||
"path": r.URL.Path,
|
||||
"method": r.Method,
|
||||
"query": r.URL.RawQuery,
|
||||
"received": time.Now().Unix(),
|
||||
})
|
||||
}
|
||||
|
||||
// URL returns the backend server URL
|
||||
func (tb *TestBackend) URL() string {
|
||||
return tb.server.URL
|
||||
}
|
||||
|
||||
// Close shuts down the backend server
|
||||
func (tb *TestBackend) Close() {
|
||||
tb.server.Close()
|
||||
}
|
||||
|
||||
// RequestCount returns the number of requests received
|
||||
func (tb *TestBackend) RequestCount() int64 {
|
||||
return atomic.LoadInt64(&tb.requestCount)
|
||||
}
|
||||
|
||||
// LastRequest returns the most recent request
|
||||
func (tb *TestBackend) LastRequest() *RequestLogEntry {
|
||||
tb.mu.Lock()
|
||||
defer tb.mu.Unlock()
|
||||
if len(tb.requestLog) == 0 {
|
||||
return nil
|
||||
}
|
||||
return &tb.requestLog[len(tb.requestLog)-1]
|
||||
}
|
||||
|
||||
// AllRequests returns all logged requests
|
||||
func (tb *TestBackend) AllRequests() []RequestLogEntry {
|
||||
tb.mu.Lock()
|
||||
defer tb.mu.Unlock()
|
||||
result := make([]RequestLogEntry, len(tb.requestLog))
|
||||
copy(result, tb.requestLog)
|
||||
return result
|
||||
}
|
||||
|
||||
// ============== HTTP Client Helpers ==============
|
||||
|
||||
// HTTPClient is a configured HTTP client for testing
|
||||
type HTTPClient struct {
|
||||
client *http.Client
|
||||
baseURL string
|
||||
}
|
||||
|
||||
// NewHTTPClient creates a new test HTTP client
|
||||
func NewHTTPClient(baseURL string) *HTTPClient {
|
||||
return &HTTPClient{
|
||||
client: &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse // Don't follow redirects
|
||||
},
|
||||
},
|
||||
baseURL: baseURL,
|
||||
}
|
||||
}
|
||||
|
||||
// Get performs a GET request
|
||||
func (c *HTTPClient) Get(path string, headers map[string]string) (*http.Response, error) {
|
||||
return c.Do("GET", path, nil, headers)
|
||||
}
|
||||
|
||||
// Post performs a POST request
|
||||
func (c *HTTPClient) Post(path string, body []byte, headers map[string]string) (*http.Response, error) {
|
||||
return c.Do("POST", path, body, headers)
|
||||
}
|
||||
|
||||
// Do performs an HTTP request
|
||||
func (c *HTTPClient) Do(method, path string, body []byte, headers map[string]string) (*http.Response, error) {
|
||||
var bodyReader io.Reader
|
||||
if body != nil {
|
||||
bodyReader = bytes.NewReader(body)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(method, c.baseURL+path, bodyReader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
return c.client.Do(req)
|
||||
}
|
||||
|
||||
// GetJSON performs GET and decodes JSON response
|
||||
func (c *HTTPClient) GetJSON(path string, result interface{}) (*http.Response, error) {
|
||||
resp, err := c.Get(path, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if err := json.NewDecoder(resp.Body).Decode(result); err != nil {
|
||||
return resp, err
|
||||
}
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// ============== File System Helpers ==============
|
||||
|
||||
// CreateTempDir creates a temporary directory for static files
|
||||
func CreateTempDir(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir, err := os.MkdirTemp("", "konduktor-test-*")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create temp dir: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { os.RemoveAll(dir) })
|
||||
return dir
|
||||
}
|
||||
|
||||
// CreateTempFile creates a temporary file with given content
|
||||
func CreateTempFile(t *testing.T, dir, name, content string) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(dir, name)
|
||||
|
||||
// Create parent directories if needed
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
|
||||
t.Fatalf("Failed to create directories: %v", err)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
|
||||
t.Fatalf("Failed to write file: %v", err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// ============== Assertion Helpers ==============
|
||||
|
||||
// AssertStatus checks if response has expected status code
|
||||
func AssertStatus(t *testing.T, resp *http.Response, expected int) {
|
||||
t.Helper()
|
||||
if resp.StatusCode != expected {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
t.Errorf("Expected status %d, got %d. Body: %s", expected, resp.StatusCode, string(body))
|
||||
}
|
||||
}
|
||||
|
||||
// AssertHeader checks if response has expected header value
|
||||
func AssertHeader(t *testing.T, resp *http.Response, header, expected string) {
|
||||
t.Helper()
|
||||
actual := resp.Header.Get(header)
|
||||
if actual != expected {
|
||||
t.Errorf("Expected header %s=%q, got %q", header, expected, actual)
|
||||
}
|
||||
}
|
||||
|
||||
// AssertHeaderContains checks if header contains substring
|
||||
func AssertHeaderContains(t *testing.T, resp *http.Response, header, substring string) {
|
||||
t.Helper()
|
||||
actual := resp.Header.Get(header)
|
||||
if actual == "" || !contains(actual, substring) {
|
||||
t.Errorf("Expected header %s to contain %q, got %q", header, substring, actual)
|
||||
}
|
||||
}
|
||||
|
||||
// AssertJSONField checks if JSON response has expected field value
|
||||
func AssertJSONField(t *testing.T, body []byte, field string, expected interface{}) {
|
||||
t.Helper()
|
||||
var data map[string]interface{}
|
||||
if err := json.Unmarshal(body, &data); err != nil {
|
||||
t.Fatalf("Failed to parse JSON: %v", err)
|
||||
}
|
||||
|
||||
actual, ok := data[field]
|
||||
if !ok {
|
||||
t.Errorf("Field %q not found in JSON", field)
|
||||
return
|
||||
}
|
||||
|
||||
if actual != expected {
|
||||
t.Errorf("Expected %s=%v, got %v", field, expected, actual)
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, substr string) bool {
|
||||
return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsAt(s, substr, 0))
|
||||
}
|
||||
|
||||
func containsAt(s, substr string, start int) bool {
|
||||
for i := start; i <= len(s)-len(substr); i++ {
|
||||
if s[i:i+len(substr)] == substr {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ReadBody reads and returns response body
|
||||
func ReadBody(t *testing.T, resp *http.Response) []byte {
|
||||
t.Helper()
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to read body: %v", err)
|
||||
}
|
||||
return body
|
||||
}
|
||||
@@ -0,0 +1,562 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/extension"
|
||||
"github.com/konduktor/konduktor/internal/logging"
|
||||
)
|
||||
|
||||
// createTestLogger creates a logger for tests
|
||||
func createTestLogger(t *testing.T) *logging.Logger {
|
||||
t.Helper()
|
||||
logger, err := logging.New(logging.Config{Level: "DEBUG"})
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create logger: %v", err)
|
||||
}
|
||||
return logger
|
||||
}
|
||||
|
||||
// ============== Basic Reverse Proxy Tests ==============
|
||||
|
||||
func TestReverseProxy_BasicGET(t *testing.T) {
|
||||
// Start backend server
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"message": "Hello from backend",
|
||||
"path": r.URL.Path,
|
||||
"method": r.Method,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
// Create routing extension with proxy to backend
|
||||
logger := createTestLogger(t)
|
||||
routingExt, err := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create routing extension: %v", err)
|
||||
}
|
||||
|
||||
// Start Konduktor server
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
// Make request through Konduktor
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/api/test", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Verify response
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
var result map[string]interface{}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
t.Fatalf("Failed to decode response: %v", err)
|
||||
}
|
||||
|
||||
if result["message"] != "Hello from backend" {
|
||||
t.Errorf("Unexpected message: %v", result["message"])
|
||||
}
|
||||
|
||||
if result["path"] != "/api/test" {
|
||||
t.Errorf("Expected path /api/test, got %v", result["path"])
|
||||
}
|
||||
|
||||
// Verify backend received request
|
||||
if backend.RequestCount() != 1 {
|
||||
t.Errorf("Expected 1 backend request, got %d", backend.RequestCount())
|
||||
}
|
||||
}
|
||||
|
||||
func TestReverseProxy_POST(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
var body map[string]interface{}
|
||||
json.NewDecoder(r.Body).Decode(&body)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"received": body,
|
||||
"method": r.Method,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
body := []byte(`{"name":"test","value":123}`)
|
||||
resp, err := client.Post("/api/data", body, map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
var result map[string]interface{}
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
if result["method"] != "POST" {
|
||||
t.Errorf("Expected method POST, got %v", result["method"])
|
||||
}
|
||||
|
||||
received := result["received"].(map[string]interface{})
|
||||
if received["name"] != "test" {
|
||||
t.Errorf("Expected name 'test', got %v", received["name"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Exact Match Routes ==============
|
||||
|
||||
func TestReverseProxy_ExactMatchRoute(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"endpoint": "version",
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
// Exact match - should use backend URL as-is
|
||||
"=/api/version": map[string]interface{}{
|
||||
"proxy_pass": backend.URL() + "/releases/latest",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Test exact match route
|
||||
resp, err := client.Get("/api/version", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
lastReq := backend.LastRequest()
|
||||
if lastReq == nil {
|
||||
t.Fatal("No request received by backend")
|
||||
}
|
||||
|
||||
// For exact match, the target path should be used as-is (IgnoreRequestPath=true)
|
||||
if lastReq.Path != "/releases/latest" {
|
||||
t.Errorf("Expected backend path /releases/latest, got %s", lastReq.Path)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Regex Routes with Parameters ==============
|
||||
|
||||
func TestReverseProxy_RegexRouteWithParams(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
// Regex with named group
|
||||
"~^/api/users/(?P<id>\\d+)$": map[string]interface{}{
|
||||
"proxy_pass": backend.URL() + "/v2/users/{id}",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Test regex route with parameter
|
||||
resp, err := client.Get("/api/users/42", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
lastReq := backend.LastRequest()
|
||||
if lastReq == nil {
|
||||
t.Fatal("No request received by backend")
|
||||
}
|
||||
|
||||
// Parameter {id} should be substituted
|
||||
if lastReq.Path != "/v2/users/42" {
|
||||
t.Errorf("Expected backend path /v2/users/42, got %s", lastReq.Path)
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Header Forwarding ==============
|
||||
|
||||
func TestReverseProxy_HeaderForwarding(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"x-forwarded-for": r.Header.Get("X-Forwarded-For"),
|
||||
"x-real-ip": r.Header.Get("X-Real-IP"),
|
||||
"x-custom": r.Header.Get("X-Custom"),
|
||||
"x-forwarded-host": r.Header.Get("X-Forwarded-Host"),
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
"headers": []interface{}{
|
||||
"X-Forwarded-For: $remote_addr",
|
||||
"X-Real-IP: $remote_addr",
|
||||
},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/test", map[string]string{
|
||||
"X-Custom": "custom-value",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var result map[string]string
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
// X-Custom should be forwarded
|
||||
if result["x-custom"] != "custom-value" {
|
||||
t.Errorf("Expected X-Custom header to be forwarded, got %v", result["x-custom"])
|
||||
}
|
||||
|
||||
// X-Forwarded-For should be set (will contain 127.0.0.1)
|
||||
if result["x-forwarded-for"] == "" {
|
||||
t.Error("Expected X-Forwarded-For header to be set")
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Query String ==============
|
||||
|
||||
func TestReverseProxy_QueryStringPreservation(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"query": r.URL.RawQuery,
|
||||
"foo": r.URL.Query().Get("foo"),
|
||||
"bar": r.URL.Query().Get("bar"),
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/search?foo=hello&bar=world", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var result map[string]string
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
if result["foo"] != "hello" {
|
||||
t.Errorf("Expected foo=hello, got %v", result["foo"])
|
||||
}
|
||||
|
||||
if result["bar"] != "world" {
|
||||
t.Errorf("Expected bar=world, got %v", result["bar"])
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Error Handling ==============
|
||||
|
||||
func TestReverseProxy_BackendUnavailable(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
// Non-existent backend
|
||||
"proxy_pass": "http://127.0.0.1:59999",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/test", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Should return 502 Bad Gateway
|
||||
AssertStatus(t, resp, http.StatusBadGateway)
|
||||
}
|
||||
|
||||
func TestReverseProxy_BackendTimeout(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Simulate slow backend
|
||||
time.Sleep(3 * time.Second)
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
"timeout": 0.5, // 500ms timeout
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get("/slow", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Should return 504 Gateway Timeout
|
||||
AssertStatus(t, resp, http.StatusGatewayTimeout)
|
||||
}
|
||||
|
||||
// ============== HTTP Methods ==============
|
||||
|
||||
func TestReverseProxy_AllMethods(t *testing.T) {
|
||||
methods := []string{"GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"}
|
||||
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"method": r.Method,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
for _, method := range methods {
|
||||
t.Run(method, func(t *testing.T) {
|
||||
resp, err := client.Do(method, "/resource", nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
if method != "HEAD" {
|
||||
var result map[string]string
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
if result["method"] != method {
|
||||
t.Errorf("Expected method %s, got %v", method, result["method"])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Large Bodies ==============
|
||||
|
||||
func TestReverseProxy_LargeRequestBody(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
json.NewEncoder(w).Encode(map[string]int{
|
||||
"received": len(body),
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// 1MB body
|
||||
largeBody := []byte(strings.Repeat("x", 1024*1024))
|
||||
resp, err := client.Post("/upload", largeBody, map[string]string{
|
||||
"Content-Type": "application/octet-stream",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
}
|
||||
|
||||
// ============== Concurrent Requests ==============
|
||||
|
||||
func TestReverseProxy_ConcurrentRequests(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Small delay to simulate work
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
const numRequests = 50
|
||||
results := make(chan error, numRequests)
|
||||
|
||||
for i := 0; i < numRequests; i++ {
|
||||
go func(n int) {
|
||||
client := NewHTTPClient(server.URL)
|
||||
resp, err := client.Get(fmt.Sprintf("/concurrent/%d", n), nil)
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
results <- fmt.Errorf("unexpected status: %d", resp.StatusCode)
|
||||
return
|
||||
}
|
||||
results <- nil
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Collect results
|
||||
var errors []error
|
||||
for i := 0; i < numRequests; i++ {
|
||||
if err := <-results; err != nil {
|
||||
errors = append(errors, err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(errors) > 0 {
|
||||
t.Errorf("Got %d errors in concurrent requests: %v", len(errors), errors[:min(5, len(errors))])
|
||||
}
|
||||
|
||||
// Verify all requests reached backend
|
||||
if backend.RequestCount() != numRequests {
|
||||
t.Errorf("Expected %d backend requests, got %d", numRequests, backend.RequestCount())
|
||||
}
|
||||
}
|
||||
|
||||
func min(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,494 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/konduktor/konduktor/internal/extension"
|
||||
)
|
||||
|
||||
// ============== Route Priority Tests ==============
|
||||
|
||||
func TestRouting_ExactMatchPriority(t *testing.T) {
|
||||
// Exact match should have highest priority
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"path": r.URL.Path,
|
||||
"source": "default",
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, err := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
// Exact match - highest priority
|
||||
"=/api/status": map[string]interface{}{
|
||||
"return": "200 exact-match",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
// Regex that also matches /api/status
|
||||
"~^/api/.*": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create routing extension: %v", err)
|
||||
}
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Test exact match route - should return static response
|
||||
resp, err := client.Get("/api/status", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
body := ReadBody(t, resp)
|
||||
if string(body) != "exact-match" {
|
||||
t.Errorf("Expected 'exact-match', got %q", string(body))
|
||||
}
|
||||
|
||||
// Regex route should be used for other /api/* paths
|
||||
resp2, err := client.Get("/api/other", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp2.Body.Close()
|
||||
|
||||
AssertStatus(t, resp2, http.StatusOK)
|
||||
|
||||
// Verify it went to backend
|
||||
if backend.RequestCount() != 1 {
|
||||
t.Errorf("Expected 1 backend request, got %d", backend.RequestCount())
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Case Sensitivity Tests ==============
|
||||
|
||||
func TestRouting_CaseSensitiveRegex(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
// Case-sensitive regex (~)
|
||||
"~^/API/test$": map[string]interface{}{
|
||||
"return": "200 case-sensitive",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"return": "200 default",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Exact case match should work
|
||||
resp, err := client.Get("/API/test", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body := ReadBody(t, resp)
|
||||
if string(body) != "case-sensitive" {
|
||||
t.Errorf("Expected 'case-sensitive' for /API/test, got %q", string(body))
|
||||
}
|
||||
|
||||
// Different case should NOT match
|
||||
resp2, err := client.Get("/api/test", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp2.Body.Close()
|
||||
|
||||
body2 := ReadBody(t, resp2)
|
||||
if string(body2) != "default" {
|
||||
t.Errorf("Expected 'default' for /api/test (case mismatch), got %q", string(body2))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouting_CaseInsensitiveRegex(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
// Case-insensitive regex (~*)
|
||||
"~*^/api/test$": map[string]interface{}{
|
||||
"return": "200 case-insensitive",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"return": "200 default",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
testCases := []struct {
|
||||
path string
|
||||
expected string
|
||||
}{
|
||||
{"/api/test", "case-insensitive"},
|
||||
{"/API/test", "case-insensitive"},
|
||||
{"/Api/Test", "case-insensitive"},
|
||||
{"/API/TEST", "case-insensitive"},
|
||||
{"/api/other", "default"},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.path, func(t *testing.T) {
|
||||
resp, err := client.Get(tc.path, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body := ReadBody(t, resp)
|
||||
if string(body) != tc.expected {
|
||||
t.Errorf("Expected %q for %s, got %q", tc.expected, tc.path, string(body))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Default Route Tests ==============
|
||||
|
||||
func TestRouting_DefaultRoute(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"handler": "default",
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"=/specific": map[string]interface{}{
|
||||
"return": "200 specific",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Non-matching paths should go to default
|
||||
paths := []string{"/", "/random", "/path/to/resource", "/api/v1/users"}
|
||||
|
||||
for _, path := range paths {
|
||||
t.Run(path, func(t *testing.T) {
|
||||
resp, err := client.Get(path, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
var result map[string]string
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
if result["handler"] != "default" {
|
||||
t.Errorf("Expected default handler, got %v", result["handler"])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Return Directive Tests ==============
|
||||
|
||||
func TestRouting_ReturnDirective(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"=/health": map[string]interface{}{
|
||||
"return": "200 OK",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"=/status": map[string]interface{}{
|
||||
"return": "200 {\"status\": \"healthy\"}",
|
||||
"content_type": "application/json",
|
||||
},
|
||||
"=/forbidden": map[string]interface{}{
|
||||
"return": "404 Not Found",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"return": "200 default",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
testCases := []struct {
|
||||
path string
|
||||
expectedStatus int
|
||||
expectedBody string
|
||||
contentType string
|
||||
}{
|
||||
{"/health", 200, "OK", "text/plain"},
|
||||
{"/status", 200, `{"status": "healthy"}`, "application/json"},
|
||||
{"/forbidden", 404, "Not Found", "text/plain"},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.path, func(t *testing.T) {
|
||||
resp, err := client.Get(tc.path, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, tc.expectedStatus)
|
||||
AssertHeaderContains(t, resp, "Content-Type", tc.contentType)
|
||||
|
||||
body := ReadBody(t, resp)
|
||||
if string(body) != tc.expectedBody {
|
||||
t.Errorf("Expected body %q, got %q", tc.expectedBody, string(body))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Multiple Regex Routes Tests ==============
|
||||
|
||||
func TestRouting_MultipleRegexRoutes(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"~^/api/v1/.*": map[string]interface{}{
|
||||
"return": "200 v1",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"~^/api/v2/.*": map[string]interface{}{
|
||||
"return": "200 v2",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"~^/api/.*": map[string]interface{}{
|
||||
"return": "200 api-generic",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"return": "200 default",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
testCases := []struct {
|
||||
path string
|
||||
expected string
|
||||
}{
|
||||
{"/api/v1/users", "v1"},
|
||||
{"/api/v2/users", "v2"},
|
||||
{"/api/v3/users", "api-generic"},
|
||||
{"/other", "default"},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.path, func(t *testing.T) {
|
||||
resp, err := client.Get(tc.path, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body := ReadBody(t, resp)
|
||||
if string(body) != tc.expected {
|
||||
t.Errorf("Expected %q for %s, got %q", tc.expected, tc.path, string(body))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== Regex with Named Groups ==============
|
||||
|
||||
func TestRouting_RegexNamedGroups(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"path": r.URL.Path,
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"~^/users/(?P<userId>\\d+)/posts/(?P<postId>\\d+)$": map[string]interface{}{
|
||||
"proxy_pass": backend.URL() + "/api/v2/users/{userId}/posts/{postId}",
|
||||
},
|
||||
"~^/items/(?P<category>[a-z]+)/(?P<id>\\d+)$": map[string]interface{}{
|
||||
"proxy_pass": backend.URL() + "/catalog/{category}/item/{id}",
|
||||
},
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
testCases := []struct {
|
||||
requestPath string
|
||||
expectedPath string
|
||||
}{
|
||||
{"/users/123/posts/456", "/api/v2/users/123/posts/456"},
|
||||
{"/items/electronics/789", "/catalog/electronics/item/789"},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.requestPath, func(t *testing.T) {
|
||||
resp, err := client.Get(tc.requestPath, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusOK)
|
||||
|
||||
lastReq := backend.LastRequest()
|
||||
if lastReq == nil {
|
||||
t.Fatal("No request received by backend")
|
||||
}
|
||||
|
||||
if lastReq.Path != tc.expectedPath {
|
||||
t.Errorf("Expected backend path %s, got %s", tc.expectedPath, lastReq.Path)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ============== No Matching Route Tests ==============
|
||||
|
||||
func TestRouting_NoMatchingRoute(t *testing.T) {
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"=/specific": map[string]interface{}{
|
||||
"return": "200 specific",
|
||||
"content_type": "text/plain",
|
||||
},
|
||||
// No default route
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
// Request to non-matching path should return 404
|
||||
resp, err := client.Get("/other", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
AssertStatus(t, resp, http.StatusNotFound)
|
||||
}
|
||||
|
||||
// ============== Headers in Return Tests ==============
|
||||
|
||||
func TestRouting_CustomHeaders(t *testing.T) {
|
||||
backend := StartBackend(func(w http.ResponseWriter, r *http.Request) {
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"x-custom-header": r.Header.Get("X-Custom-Header"),
|
||||
"x-api-version": r.Header.Get("X-API-Version"),
|
||||
})
|
||||
})
|
||||
defer backend.Close()
|
||||
|
||||
logger := createTestLogger(t)
|
||||
routingExt, _ := extension.NewRoutingExtension(map[string]interface{}{
|
||||
"regex_locations": map[string]interface{}{
|
||||
"__default__": map[string]interface{}{
|
||||
"proxy_pass": backend.URL(),
|
||||
"headers": []interface{}{
|
||||
"X-Custom-Header: custom-value",
|
||||
"X-API-Version: v1",
|
||||
},
|
||||
},
|
||||
},
|
||||
}, logger)
|
||||
|
||||
server := StartTestServer(t, &ServerConfig{
|
||||
Extensions: []extension.Extension{routingExt},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
client := NewHTTPClient(server.URL)
|
||||
|
||||
resp, err := client.Get("/test", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var result map[string]string
|
||||
json.NewDecoder(resp.Body).Decode(&result)
|
||||
|
||||
if result["x-custom-header"] != "custom-value" {
|
||||
t.Errorf("Expected X-Custom-Header=custom-value, got %v", result["x-custom-header"])
|
||||
}
|
||||
|
||||
if result["x-api-version"] != "v1" {
|
||||
t.Errorf("Expected X-API-Version=v1, got %v", result["x-api-version"])
|
||||
}
|
||||
}
|
||||
Generated
+104
-4
@@ -147,14 +147,14 @@ files = [
|
||||
|
||||
[[package]]
|
||||
name = "click"
|
||||
version = "8.2.1"
|
||||
version = "8.3.1"
|
||||
description = "Composable command line interface toolkit"
|
||||
optional = false
|
||||
python-versions = ">=3.10"
|
||||
groups = ["main", "dev"]
|
||||
files = [
|
||||
{file = "click-8.2.1-py3-none-any.whl", hash = "sha256:61a3265b914e850b85317d0b3109c7f8cd35a670f963866005d6ef1d5175a12b"},
|
||||
{file = "click-8.2.1.tar.gz", hash = "sha256:27c491cc05d968d271d5a1db13e3b5a184636d9d930f148c50b038f0d0646202"},
|
||||
{file = "click-8.3.1-py3-none-any.whl", hash = "sha256:981153a64e25f12d547d3426c367a4857371575ee7ad18df2a6183ab0545b2a6"},
|
||||
{file = "click-8.3.1.tar.gz", hash = "sha256:12ff4785d337a1bb490bb7e9c2b1ee5da3112e94a8622f26a6c77f5d2fc6842a"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
@@ -602,6 +602,30 @@ MarkupSafe = ">=2.0"
|
||||
[package.extras]
|
||||
i18n = ["Babel (>=2.7)"]
|
||||
|
||||
[[package]]
|
||||
name = "markdown-it-py"
|
||||
version = "4.0.0"
|
||||
description = "Python port of markdown-it. Markdown parsing, done right!"
|
||||
optional = false
|
||||
python-versions = ">=3.10"
|
||||
groups = ["main"]
|
||||
files = [
|
||||
{file = "markdown_it_py-4.0.0-py3-none-any.whl", hash = "sha256:87327c59b172c5011896038353a81343b6754500a08cd7a4973bb48c6d578147"},
|
||||
{file = "markdown_it_py-4.0.0.tar.gz", hash = "sha256:cb0a2b4aa34f932c007117b194e945bd74e0ec24133ceb5bac59009cda1cb9f3"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
mdurl = ">=0.1,<1.0"
|
||||
|
||||
[package.extras]
|
||||
benchmarking = ["psutil", "pytest", "pytest-benchmark"]
|
||||
compare = ["commonmark (>=0.9,<1.0)", "markdown (>=3.4,<4.0)", "markdown-it-pyrs", "mistletoe (>=1.0,<2.0)", "mistune (>=3.0,<4.0)", "panflute (>=2.3,<3.0)"]
|
||||
linkify = ["linkify-it-py (>=1,<3)"]
|
||||
plugins = ["mdit-py-plugins (>=0.5.0)"]
|
||||
profiling = ["gprof2dot"]
|
||||
rtd = ["ipykernel", "jupyter_sphinx", "mdit-py-plugins (>=0.5.0)", "myst-parser", "pyyaml", "sphinx", "sphinx-book-theme (>=1.0,<2.0)", "sphinx-copybutton", "sphinx-design"]
|
||||
testing = ["coverage", "pytest", "pytest-cov", "pytest-regressions", "requests"]
|
||||
|
||||
[[package]]
|
||||
name = "markupsafe"
|
||||
version = "3.0.3"
|
||||
@@ -714,6 +738,18 @@ files = [
|
||||
{file = "mccabe-0.7.0.tar.gz", hash = "sha256:348e0240c33b60bbdf4e523192ef919f28cb2c3d7d5c7794f74009290f236325"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mdurl"
|
||||
version = "0.1.2"
|
||||
description = "Markdown URL utilities"
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
groups = ["main"]
|
||||
files = [
|
||||
{file = "mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8"},
|
||||
{file = "mdurl-0.1.2.tar.gz", hash = "sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mypy"
|
||||
version = "1.17.1"
|
||||
@@ -843,6 +879,39 @@ files = [
|
||||
dev = ["pre-commit", "tox"]
|
||||
testing = ["coverage", "pytest", "pytest-benchmark"]
|
||||
|
||||
[[package]]
|
||||
name = "psutil"
|
||||
version = "7.1.3"
|
||||
description = "Cross-platform lib for process and system monitoring."
|
||||
optional = false
|
||||
python-versions = ">=3.6"
|
||||
groups = ["main"]
|
||||
files = [
|
||||
{file = "psutil-7.1.3-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:0005da714eee687b4b8decd3d6cc7c6db36215c9e74e5ad2264b90c3df7d92dc"},
|
||||
{file = "psutil-7.1.3-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:19644c85dcb987e35eeeaefdc3915d059dac7bd1167cdcdbf27e0ce2df0c08c0"},
|
||||
{file = "psutil-7.1.3-cp313-cp313t-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:95ef04cf2e5ba0ab9eaafc4a11eaae91b44f4ef5541acd2ee91d9108d00d59a7"},
|
||||
{file = "psutil-7.1.3-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1068c303be3a72f8e18e412c5b2a8f6d31750fb152f9cb106b54090296c9d251"},
|
||||
{file = "psutil-7.1.3-cp313-cp313t-win_amd64.whl", hash = "sha256:18349c5c24b06ac5612c0428ec2a0331c26443d259e2a0144a9b24b4395b58fa"},
|
||||
{file = "psutil-7.1.3-cp313-cp313t-win_arm64.whl", hash = "sha256:c525ffa774fe4496282fb0b1187725793de3e7c6b29e41562733cae9ada151ee"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:b403da1df4d6d43973dc004d19cee3b848e998ae3154cc8097d139b77156c353"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:ad81425efc5e75da3f39b3e636293360ad8d0b49bed7df824c79764fb4ba9b8b"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8f33a3702e167783a9213db10ad29650ebf383946e91bc77f28a5eb083496bc9"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:fac9cd332c67f4422504297889da5ab7e05fd11e3c4392140f7370f4208ded1f"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-win_amd64.whl", hash = "sha256:3792983e23b69843aea49c8f5b8f115572c5ab64c153bada5270086a2123c7e7"},
|
||||
{file = "psutil-7.1.3-cp314-cp314t-win_arm64.whl", hash = "sha256:31d77fcedb7529f27bb3a0472bea9334349f9a04160e8e6e5020f22c59893264"},
|
||||
{file = "psutil-7.1.3-cp36-abi3-macosx_10_9_x86_64.whl", hash = "sha256:2bdbcd0e58ca14996a42adf3621a6244f1bb2e2e528886959c72cf1e326677ab"},
|
||||
{file = "psutil-7.1.3-cp36-abi3-macosx_11_0_arm64.whl", hash = "sha256:bc31fa00f1fbc3c3802141eede66f3a2d51d89716a194bf2cd6fc68310a19880"},
|
||||
{file = "psutil-7.1.3-cp36-abi3-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3bb428f9f05c1225a558f53e30ccbad9930b11c3fc206836242de1091d3e7dd3"},
|
||||
{file = "psutil-7.1.3-cp36-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:56d974e02ca2c8eb4812c3f76c30e28836fffc311d55d979f1465c1feeb2b68b"},
|
||||
{file = "psutil-7.1.3-cp37-abi3-win_amd64.whl", hash = "sha256:f39c2c19fe824b47484b96f9692932248a54c43799a84282cfe58d05a6449efd"},
|
||||
{file = "psutil-7.1.3-cp37-abi3-win_arm64.whl", hash = "sha256:bd0d69cee829226a761e92f28140bec9a5ee9d5b4fb4b0cc589068dbfff559b1"},
|
||||
{file = "psutil-7.1.3.tar.gz", hash = "sha256:6c86281738d77335af7aec228328e944b30930899ea760ecf33a4dba66be5e74"},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
dev = ["abi3audit", "black", "check-manifest", "colorama ; os_name == \"nt\"", "coverage", "packaging", "pylint", "pyperf", "pypinfo", "pyreadline ; os_name == \"nt\"", "pytest", "pytest-cov", "pytest-instafail", "pytest-subtests", "pytest-xdist", "pywin32 ; os_name == \"nt\" and platform_python_implementation != \"PyPy\"", "requests", "rstcheck", "ruff", "setuptools", "sphinx", "sphinx_rtd_theme", "toml-sort", "twine", "validate-pyproject[all]", "virtualenv", "vulture", "wheel", "wheel ; os_name == \"nt\" and platform_python_implementation != \"PyPy\"", "wmi ; os_name == \"nt\" and platform_python_implementation != \"PyPy\""]
|
||||
test = ["pytest", "pytest-instafail", "pytest-subtests", "pytest-xdist", "pywin32 ; os_name == \"nt\" and platform_python_implementation != \"PyPy\"", "setuptools", "wheel ; os_name == \"nt\" and platform_python_implementation != \"PyPy\"", "wmi ; os_name == \"nt\" and platform_python_implementation != \"PyPy\""]
|
||||
|
||||
[[package]]
|
||||
name = "pycodestyle"
|
||||
version = "2.14.0"
|
||||
@@ -1180,6 +1249,25 @@ files = [
|
||||
{file = "pyyaml-6.0.2.tar.gz", hash = "sha256:d584d9ec91ad65861cc08d42e834324ef890a082e591037abe114850ff7bbc3e"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rich"
|
||||
version = "14.2.0"
|
||||
description = "Render rich text, tables, progress bars, syntax highlighting, markdown and more to the terminal"
|
||||
optional = false
|
||||
python-versions = ">=3.8.0"
|
||||
groups = ["main"]
|
||||
files = [
|
||||
{file = "rich-14.2.0-py3-none-any.whl", hash = "sha256:76bc51fe2e57d2b1be1f96c524b890b816e334ab4c1e45888799bfaab0021edd"},
|
||||
{file = "rich-14.2.0.tar.gz", hash = "sha256:73ff50c7c0c1c77c8243079283f4edb376f0f6442433aecb8ce7e6d0b92d1fe4"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
markdown-it-py = ">=2.2.0"
|
||||
pygments = ">=2.13.0,<3.0.0"
|
||||
|
||||
[package.extras]
|
||||
jupyter = ["ipywidgets (>=7.5.1,<9)"]
|
||||
|
||||
[[package]]
|
||||
name = "setuptools"
|
||||
version = "80.9.0"
|
||||
@@ -1261,6 +1349,18 @@ files = [
|
||||
{file = "structlog-25.4.0.tar.gz", hash = "sha256:186cd1b0a8ae762e29417095664adf1d6a31702160a46dacb7796ea82f7409e4"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "types-psutil"
|
||||
version = "7.1.3.20251202"
|
||||
description = "Typing stubs for psutil"
|
||||
optional = false
|
||||
python-versions = ">=3.9"
|
||||
groups = ["dev"]
|
||||
files = [
|
||||
{file = "types_psutil-7.1.3.20251202-py3-none-any.whl", hash = "sha256:39bfc44780de7ab686c65169e36a7969db09e7f39d92de643b55789292953400"},
|
||||
{file = "types_psutil-7.1.3.20251202.tar.gz", hash = "sha256:5cfecaced7c486fb3995bb290eab45043d697a261718aca01b9b340d1ab7968a"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "types-pyyaml"
|
||||
version = "6.0.12.20250822"
|
||||
@@ -1621,4 +1721,4 @@ wsgi = ["a2wsgi"]
|
||||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.12"
|
||||
content-hash = "32ebf260f6792987cb4236fe29ad3329374e063504d507b5a0319684e24a30a8"
|
||||
content-hash = "653d7b992e2bb133abde2e8b1c44265e948ed90487ab3f2670429510a8aa0683"
|
||||
|
||||
+7
-2
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "pyserve"
|
||||
version = "0.9.1"
|
||||
description = "Simple HTTP Web server written in Python"
|
||||
version = "0.9.10"
|
||||
description = "Python Application Orchestrator & HTTP Server - unified gateway for multiple Python web apps"
|
||||
authors = [
|
||||
{name = "Илья Глазунов",email = "i.glazunov@sapiens.solutions"}
|
||||
]
|
||||
@@ -15,10 +15,14 @@ dependencies = [
|
||||
"types-pyyaml (>=6.0.12.20250822,<7.0.0.0)",
|
||||
"structlog (>=25.4.0,<26.0.0)",
|
||||
"httpx (>=0.27.0,<0.28.0)",
|
||||
"click (>=8.0)",
|
||||
"rich (>=13.0)",
|
||||
"psutil (>=5.9)",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
pyserve = "pyserve.cli:main"
|
||||
pyservectl = "pyserve.ctl:main"
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
@@ -97,4 +101,5 @@ flake8 = "^7.3.0"
|
||||
pytest-asyncio = "^1.3.0"
|
||||
cython = "^3.0.0"
|
||||
setuptools = "^80.0.0"
|
||||
types-psutil = "^7.1.3.20251202"
|
||||
|
||||
|
||||
+19
-2
@@ -2,7 +2,7 @@
|
||||
PyServe - HTTP web server written on Python
|
||||
"""
|
||||
|
||||
__version__ = "0.9.0"
|
||||
__version__ = "0.10.0"
|
||||
__author__ = "Ilya Glazunov"
|
||||
|
||||
from .asgi_mount import (
|
||||
@@ -15,13 +15,22 @@ from .asgi_mount import (
|
||||
create_starlette_app,
|
||||
)
|
||||
from .config import Config
|
||||
from .process_manager import (
|
||||
ProcessConfig,
|
||||
ProcessInfo,
|
||||
ProcessManager,
|
||||
ProcessState,
|
||||
get_process_manager,
|
||||
init_process_manager,
|
||||
shutdown_process_manager,
|
||||
)
|
||||
from .server import PyServeServer
|
||||
|
||||
__all__ = [
|
||||
"PyServeServer",
|
||||
"Config",
|
||||
"__version__",
|
||||
# ASGI mounting
|
||||
# ASGI mounting (in-process)
|
||||
"ASGIAppLoader",
|
||||
"ASGIMountManager",
|
||||
"MountedApp",
|
||||
@@ -29,4 +38,12 @@ __all__ = [
|
||||
"create_flask_app",
|
||||
"create_django_app",
|
||||
"create_starlette_app",
|
||||
# Process orchestration (multi-process)
|
||||
"ProcessManager",
|
||||
"ProcessConfig",
|
||||
"ProcessInfo",
|
||||
"ProcessState",
|
||||
"get_process_manager",
|
||||
"init_process_manager",
|
||||
"shutdown_process_manager",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
"""
|
||||
WSGI Wrapper Module for Process Orchestration.
|
||||
|
||||
This module provides a wrapper that allows WSGI applications to be run
|
||||
via uvicorn by wrapping them with a2wsgi.
|
||||
|
||||
The WSGI app path is passed via environment variables:
|
||||
- PYSERVE_WSGI_APP: The app path (e.g., "myapp:app" or "myapp.main:create_app")
|
||||
- PYSERVE_WSGI_FACTORY: "1" if the app path points to a factory function
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import os
|
||||
from typing import Any, Callable, Optional, Type
|
||||
|
||||
WSGIMiddlewareType = Optional[Type[Any]]
|
||||
WSGI_ADAPTER: Optional[str] = None
|
||||
WSGIMiddleware: WSGIMiddlewareType = None
|
||||
|
||||
try:
|
||||
from a2wsgi import WSGIMiddleware as _A2WSGIMiddleware
|
||||
|
||||
WSGIMiddleware = _A2WSGIMiddleware
|
||||
WSGI_ADAPTER = "a2wsgi"
|
||||
except ImportError:
|
||||
try:
|
||||
from asgiref.wsgi import WsgiToAsgi as _AsgirefMiddleware
|
||||
|
||||
WSGIMiddleware = _AsgirefMiddleware
|
||||
WSGI_ADAPTER = "asgiref"
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
def _load_wsgi_app() -> Callable[..., Any]:
|
||||
app_path = os.environ.get("PYSERVE_WSGI_APP")
|
||||
is_factory = os.environ.get("PYSERVE_WSGI_FACTORY", "0") == "1"
|
||||
|
||||
if not app_path:
|
||||
raise RuntimeError("PYSERVE_WSGI_APP environment variable not set. " "This module should only be used by PyServe process orchestration.")
|
||||
|
||||
if ":" in app_path:
|
||||
module_name, attr_name = app_path.rsplit(":", 1)
|
||||
else:
|
||||
module_name = app_path
|
||||
attr_name = "app"
|
||||
|
||||
try:
|
||||
module = importlib.import_module(module_name)
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"Failed to import WSGI module '{module_name}': {e}")
|
||||
|
||||
try:
|
||||
app_or_factory = getattr(module, attr_name)
|
||||
except AttributeError:
|
||||
raise RuntimeError(f"Module '{module_name}' has no attribute '{attr_name}'")
|
||||
|
||||
if is_factory:
|
||||
result: Callable[..., Any] = app_or_factory()
|
||||
return result
|
||||
loaded_app: Callable[..., Any] = app_or_factory
|
||||
return loaded_app
|
||||
|
||||
|
||||
def _create_asgi_app() -> Any:
|
||||
if WSGIMiddleware is None:
|
||||
raise RuntimeError("No WSGI adapter available. " "Install a2wsgi (recommended) or asgiref: pip install a2wsgi")
|
||||
|
||||
wsgi_app = _load_wsgi_app()
|
||||
return WSGIMiddleware(wsgi_app)
|
||||
|
||||
|
||||
app = _create_asgi_app()
|
||||
+33
-5
@@ -1,3 +1,10 @@
|
||||
"""
|
||||
PyServe CLI - Server entry point
|
||||
|
||||
Simple CLI for running the PyServe HTTP server.
|
||||
For service management, use pyservectl.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
@@ -9,12 +16,33 @@ def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="PyServe - HTTP web server",
|
||||
prog="pyserve",
|
||||
epilog="For service management (start/stop/restart/logs), use: pyservectl",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-c",
|
||||
"--config",
|
||||
default="config.yaml",
|
||||
help="Path to configuration file (default: config.yaml)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--host",
|
||||
help="Host to bind the server to",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
help="Port to bind the server to",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug",
|
||||
action="store_true",
|
||||
help="Enable debug mode",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--version",
|
||||
action="version",
|
||||
version=f"%(prog)s {__version__}",
|
||||
)
|
||||
parser.add_argument("-c", "--config", default="config.yaml", help="Path to configuration file (default: config.yaml)")
|
||||
parser.add_argument("--host", help="Host to bind the server to")
|
||||
parser.add_argument("--port", type=int, help="Port to bind the server to")
|
||||
parser.add_argument("--debug", action="store_true", help="Enable debug mode")
|
||||
parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
"""
|
||||
PyServeCtl - Service management CLI
|
||||
|
||||
Docker-compose-like tool for managing PyServe services.
|
||||
|
||||
Usage:
|
||||
pyservectl [OPTIONS] COMMAND [ARGS]...
|
||||
|
||||
Commands:
|
||||
init Initialize a new project
|
||||
config Configuration management
|
||||
up Start all services
|
||||
down Stop all services
|
||||
start Start specific services
|
||||
stop Stop specific services
|
||||
restart Restart services
|
||||
ps Show service status
|
||||
logs View service logs
|
||||
top Live monitoring dashboard
|
||||
health Check service health
|
||||
scale Scale services
|
||||
"""
|
||||
|
||||
from .main import cli, main
|
||||
|
||||
__all__ = ["cli", "main"]
|
||||
@@ -0,0 +1,93 @@
|
||||
"""
|
||||
PyServe Daemon Process
|
||||
|
||||
Runs pyserve services in background mode.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import FrameType
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="PyServe Daemon")
|
||||
parser.add_argument("--config", required=True, help="Configuration file path")
|
||||
parser.add_argument("--state-dir", required=True, help="State directory path")
|
||||
parser.add_argument("--services", default=None, help="Comma-separated list of services")
|
||||
parser.add_argument("--scale", action="append", default=[], help="Scale overrides (name=workers)")
|
||||
parser.add_argument("--force-recreate", action="store_true", help="Force recreate services")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
config_path = Path(args.config)
|
||||
state_dir = Path(args.state_dir)
|
||||
|
||||
services = args.services.split(",") if args.services else None
|
||||
|
||||
scale_map = {}
|
||||
for scale in args.scale:
|
||||
name, workers = scale.split("=")
|
||||
scale_map[name] = int(workers)
|
||||
|
||||
from ..config import Config
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
|
||||
from .state import StateManager
|
||||
|
||||
state_manager = StateManager(state_dir)
|
||||
|
||||
log_file = state_dir / "logs" / "daemon.log"
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
handlers=[
|
||||
logging.FileHandler(log_file),
|
||||
],
|
||||
)
|
||||
|
||||
logger = logging.getLogger("pyserve.daemon")
|
||||
|
||||
pid_file = state_dir / "pyserve.pid"
|
||||
pid_file.write_text(str(os.getpid()))
|
||||
|
||||
logger.info(f"Starting daemon with PID {os.getpid()}")
|
||||
|
||||
from ._runner import ServiceRunner
|
||||
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
def signal_handler(signum: int, frame: Optional[FrameType]) -> None:
|
||||
logger.info(f"Received signal {signum}, shutting down...")
|
||||
runner.stop()
|
||||
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
|
||||
try:
|
||||
asyncio.run(
|
||||
runner.start(
|
||||
services=services,
|
||||
scale_map=scale_map,
|
||||
force_recreate=args.force_recreate,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Daemon error: {e}")
|
||||
sys.exit(1)
|
||||
finally:
|
||||
if pid_file.exists():
|
||||
pid_file.unlink()
|
||||
logger.info("Daemon stopped")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,389 @@
|
||||
"""
|
||||
PyServe Service Runner
|
||||
|
||||
Handles starting, stopping, and managing services.
|
||||
Integrates with ProcessManager for actual process management.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from ..config import Config
|
||||
from ..process_manager import ProcessConfig, ProcessManager, ProcessState
|
||||
from .state import StateManager
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServiceDefinition:
|
||||
name: str
|
||||
path: str
|
||||
app_path: str
|
||||
app_type: str = "asgi"
|
||||
module_path: Optional[str] = None
|
||||
workers: int = 1
|
||||
health_check_path: str = "/health"
|
||||
health_check_interval: float = 10.0
|
||||
health_check_timeout: float = 5.0
|
||||
health_check_retries: int = 3
|
||||
max_restart_count: int = 5
|
||||
restart_delay: float = 1.0
|
||||
shutdown_timeout: float = 30.0
|
||||
strip_path: bool = True
|
||||
env: Dict[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
class ServiceRunner:
|
||||
def __init__(self, config: Config, state_manager: StateManager):
|
||||
self.config = config
|
||||
self.state_manager = state_manager
|
||||
self._process_manager: Optional[ProcessManager] = None
|
||||
self._services: Dict[str, ServiceDefinition] = {}
|
||||
self._running = False
|
||||
|
||||
self._parse_services()
|
||||
|
||||
def _parse_services(self) -> None:
|
||||
for ext in self.config.extensions:
|
||||
if ext.type == "process_orchestration":
|
||||
apps = ext.config.get("apps", [])
|
||||
for app_config in apps:
|
||||
service = ServiceDefinition(
|
||||
name=app_config.get("name", "unnamed"),
|
||||
path=app_config.get("path", "/"),
|
||||
app_path=app_config.get("app_path", ""),
|
||||
app_type=app_config.get("app_type", "asgi"),
|
||||
module_path=app_config.get("module_path"),
|
||||
workers=app_config.get("workers", 1),
|
||||
health_check_path=app_config.get("health_check_path", "/health"),
|
||||
health_check_interval=app_config.get("health_check_interval", 10.0),
|
||||
health_check_timeout=app_config.get("health_check_timeout", 5.0),
|
||||
health_check_retries=app_config.get("health_check_retries", 3),
|
||||
max_restart_count=app_config.get("max_restart_count", 5),
|
||||
restart_delay=app_config.get("restart_delay", 1.0),
|
||||
shutdown_timeout=app_config.get("shutdown_timeout", 30.0),
|
||||
strip_path=app_config.get("strip_path", True),
|
||||
env=app_config.get("env", {}),
|
||||
)
|
||||
self._services[service.name] = service
|
||||
|
||||
def get_services(self) -> Dict[str, ServiceDefinition]:
|
||||
return self._services.copy()
|
||||
|
||||
def get_service(self, name: str) -> Optional[ServiceDefinition]:
|
||||
return self._services.get(name)
|
||||
|
||||
async def start(
|
||||
self,
|
||||
services: Optional[List[str]] = None,
|
||||
scale_map: Optional[Dict[str, int]] = None,
|
||||
force_recreate: bool = False,
|
||||
wait_healthy: bool = False,
|
||||
timeout: int = 60,
|
||||
) -> None:
|
||||
from .output import console, print_error, print_info, print_success
|
||||
|
||||
scale_map = scale_map or {}
|
||||
|
||||
target_services = services or list(self._services.keys())
|
||||
|
||||
if not target_services:
|
||||
print_info("No services configured. Add services to your config.yaml")
|
||||
return
|
||||
|
||||
for name in target_services:
|
||||
if name not in self._services:
|
||||
print_error(f"Service '{name}' not found in configuration")
|
||||
return
|
||||
|
||||
port_range = (9000, 9999)
|
||||
for ext in self.config.extensions:
|
||||
if ext.type == "process_orchestration":
|
||||
port_range = tuple(ext.config.get("port_range", [9000, 9999]))
|
||||
break
|
||||
|
||||
self._process_manager = ProcessManager(
|
||||
port_range=port_range,
|
||||
health_check_enabled=True,
|
||||
)
|
||||
await self._process_manager.start()
|
||||
|
||||
self._running = True
|
||||
|
||||
for name in target_services:
|
||||
service = self._services[name]
|
||||
workers = scale_map.get(name, service.workers)
|
||||
|
||||
proc_config = ProcessConfig(
|
||||
name=name,
|
||||
app_path=service.app_path,
|
||||
app_type=service.app_type,
|
||||
workers=workers,
|
||||
module_path=service.module_path,
|
||||
health_check_enabled=True,
|
||||
health_check_path=service.health_check_path,
|
||||
health_check_interval=service.health_check_interval,
|
||||
health_check_timeout=service.health_check_timeout,
|
||||
health_check_retries=service.health_check_retries,
|
||||
max_restart_count=service.max_restart_count,
|
||||
restart_delay=service.restart_delay,
|
||||
shutdown_timeout=service.shutdown_timeout,
|
||||
env=service.env,
|
||||
)
|
||||
|
||||
try:
|
||||
await self._process_manager.register(proc_config)
|
||||
success = await self._process_manager.start_process(name)
|
||||
|
||||
if success:
|
||||
info = self._process_manager.get_process(name)
|
||||
if info:
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state="running",
|
||||
pid=info.pid,
|
||||
port=info.port,
|
||||
workers=workers,
|
||||
started_at=time.time(),
|
||||
)
|
||||
print_success(f"Started service: {name}")
|
||||
else:
|
||||
self.state_manager.update_service(name, state="failed")
|
||||
print_error(f"Failed to start service: {name}")
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error starting {name}: {e}")
|
||||
self.state_manager.update_service(name, state="failed")
|
||||
|
||||
if wait_healthy:
|
||||
print_info("Waiting for services to be healthy...")
|
||||
await self._wait_healthy(target_services, timeout)
|
||||
|
||||
console.print("\n[bold]Services running. Press Ctrl+C to stop.[/bold]\n")
|
||||
|
||||
try:
|
||||
while self._running:
|
||||
await asyncio.sleep(1)
|
||||
await self._sync_state()
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
finally:
|
||||
await self.stop_all()
|
||||
|
||||
async def _sync_state(self) -> None:
|
||||
if not self._process_manager:
|
||||
return
|
||||
|
||||
for name, info in self._process_manager.get_all_processes().items():
|
||||
state_str = info.state.value
|
||||
health_status = "healthy" if info.health_check_failures == 0 else "unhealthy"
|
||||
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state=state_str,
|
||||
pid=info.pid,
|
||||
port=info.port,
|
||||
)
|
||||
|
||||
service_state = self.state_manager.get_service(name)
|
||||
if service_state:
|
||||
service_state.health.status = health_status
|
||||
service_state.health.failures = info.health_check_failures
|
||||
self.state_manager.save()
|
||||
|
||||
async def _wait_healthy(self, services: List[str], timeout: int) -> None:
|
||||
from .output import print_info, print_warning
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
all_healthy = True
|
||||
|
||||
for name in services:
|
||||
if not self._process_manager:
|
||||
continue
|
||||
|
||||
info = self._process_manager.get_process(name)
|
||||
if not info or info.state != ProcessState.RUNNING:
|
||||
all_healthy = False
|
||||
break
|
||||
|
||||
if all_healthy:
|
||||
print_info("All services healthy")
|
||||
return
|
||||
|
||||
await asyncio.sleep(1)
|
||||
|
||||
print_warning("Timeout waiting for services to become healthy")
|
||||
|
||||
async def stop_all(self, timeout: int = 30) -> None:
|
||||
from .output import print_info
|
||||
|
||||
self._running = False
|
||||
|
||||
if self._process_manager:
|
||||
print_info("Stopping all services...")
|
||||
await self._process_manager.stop()
|
||||
self._process_manager = None
|
||||
|
||||
for name in self._services:
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state="stopped",
|
||||
pid=None,
|
||||
)
|
||||
|
||||
def stop(self) -> None:
|
||||
self._running = False
|
||||
|
||||
async def start_service(self, name: str, timeout: int = 60) -> bool:
|
||||
from .output import print_error
|
||||
|
||||
service = self._services.get(name)
|
||||
if not service:
|
||||
print_error(f"Service '{name}' not found")
|
||||
return False
|
||||
|
||||
if not self._process_manager:
|
||||
self._process_manager = ProcessManager()
|
||||
await self._process_manager.start()
|
||||
|
||||
proc_config = ProcessConfig(
|
||||
name=name,
|
||||
app_path=service.app_path,
|
||||
app_type=service.app_type,
|
||||
workers=service.workers,
|
||||
module_path=service.module_path,
|
||||
health_check_enabled=True,
|
||||
health_check_path=service.health_check_path,
|
||||
env=service.env,
|
||||
)
|
||||
|
||||
try:
|
||||
existing = self._process_manager.get_process(name)
|
||||
if not existing:
|
||||
await self._process_manager.register(proc_config)
|
||||
|
||||
success = await self._process_manager.start_process(name)
|
||||
|
||||
if success:
|
||||
info = self._process_manager.get_process(name)
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state="running",
|
||||
pid=info.pid if info else None,
|
||||
port=info.port if info else 0,
|
||||
started_at=time.time(),
|
||||
)
|
||||
|
||||
return success
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error starting {name}: {e}")
|
||||
return False
|
||||
|
||||
async def stop_service(self, name: str, timeout: int = 30, force: bool = False) -> bool:
|
||||
if not self._process_manager:
|
||||
self.state_manager.update_service(name, state="stopped", pid=None)
|
||||
return True
|
||||
|
||||
try:
|
||||
success = await self._process_manager.stop_process(name)
|
||||
|
||||
if success:
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state="stopped",
|
||||
pid=None,
|
||||
)
|
||||
|
||||
return success
|
||||
|
||||
except Exception as e:
|
||||
from .output import print_error
|
||||
|
||||
print_error(f"Error stopping {name}: {e}")
|
||||
return False
|
||||
|
||||
async def restart_service(self, name: str, timeout: int = 60) -> bool:
|
||||
if not self._process_manager:
|
||||
return False
|
||||
|
||||
try:
|
||||
self.state_manager.update_service(name, state="restarting")
|
||||
success = await self._process_manager.restart_process(name)
|
||||
|
||||
if success:
|
||||
info = self._process_manager.get_process(name)
|
||||
self.state_manager.update_service(
|
||||
name,
|
||||
state="running",
|
||||
pid=info.pid if info else None,
|
||||
port=info.port if info else 0,
|
||||
started_at=time.time(),
|
||||
)
|
||||
|
||||
return success
|
||||
|
||||
except Exception as e:
|
||||
from .output import print_error
|
||||
|
||||
print_error(f"Error restarting {name}: {e}")
|
||||
return False
|
||||
|
||||
async def scale_service(self, name: str, workers: int, timeout: int = 60, wait: bool = True) -> bool:
|
||||
# For now, this requires restart with new worker count
|
||||
# In future, could implement hot-reloading
|
||||
|
||||
service = self._services.get(name)
|
||||
if not service:
|
||||
return False
|
||||
|
||||
# Update service definition
|
||||
service.workers = workers
|
||||
|
||||
# Restart with new configuration
|
||||
return await self.restart_service(name, timeout)
|
||||
|
||||
def start_daemon(
|
||||
self,
|
||||
services: Optional[List[str]] = None,
|
||||
scale_map: Optional[Dict[str, int]] = None,
|
||||
force_recreate: bool = False,
|
||||
) -> int:
|
||||
import subprocess
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"pyserve.cli._daemon",
|
||||
"--config",
|
||||
str(self.state_manager.state_dir.parent / "config.yaml"),
|
||||
"--state-dir",
|
||||
str(self.state_manager.state_dir),
|
||||
]
|
||||
|
||||
if services:
|
||||
cmd.extend(["--services", ",".join(services)])
|
||||
|
||||
if scale_map:
|
||||
for name, workers in scale_map.items():
|
||||
cmd.extend(["--scale", f"{name}={workers}"])
|
||||
|
||||
if force_recreate:
|
||||
cmd.append("--force-recreate")
|
||||
|
||||
env = os.environ.copy()
|
||||
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
env=env,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
start_new_session=True,
|
||||
)
|
||||
|
||||
return process.pid
|
||||
@@ -0,0 +1,25 @@
|
||||
from .config import config_cmd
|
||||
from .down import down_cmd
|
||||
from .health import health_cmd
|
||||
from .init import init_cmd
|
||||
from .logs import logs_cmd
|
||||
from .scale import scale_cmd
|
||||
from .service import restart_cmd, start_cmd, stop_cmd
|
||||
from .status import ps_cmd
|
||||
from .top import top_cmd
|
||||
from .up import up_cmd
|
||||
|
||||
__all__ = [
|
||||
"init_cmd",
|
||||
"config_cmd",
|
||||
"up_cmd",
|
||||
"down_cmd",
|
||||
"start_cmd",
|
||||
"stop_cmd",
|
||||
"restart_cmd",
|
||||
"ps_cmd",
|
||||
"logs_cmd",
|
||||
"top_cmd",
|
||||
"health_cmd",
|
||||
"scale_cmd",
|
||||
]
|
||||
@@ -0,0 +1,419 @@
|
||||
"""
|
||||
pyserve config - Configuration management commands
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import click
|
||||
import yaml
|
||||
|
||||
|
||||
@click.group("config")
|
||||
def config_cmd() -> None:
|
||||
"""
|
||||
Configuration management commands.
|
||||
|
||||
\b
|
||||
Commands:
|
||||
validate Validate configuration file
|
||||
show Display current configuration
|
||||
get Get a specific configuration value
|
||||
set Set a configuration value
|
||||
diff Compare two configuration files
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
@config_cmd.command("validate")
|
||||
@click.option(
|
||||
"-c",
|
||||
"--config",
|
||||
"config_file",
|
||||
default=None,
|
||||
help="Path to configuration file",
|
||||
)
|
||||
@click.option(
|
||||
"--strict",
|
||||
is_flag=True,
|
||||
help="Enable strict validation (warn on unknown fields)",
|
||||
)
|
||||
@click.pass_obj
|
||||
def validate_cmd(ctx: Any, config_file: Optional[str], strict: bool) -> None:
|
||||
"""
|
||||
Validate a configuration file.
|
||||
|
||||
Checks for syntax errors, missing required fields, and invalid values.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve config validate
|
||||
pyserve config validate -c production.yaml
|
||||
pyserve config validate --strict
|
||||
"""
|
||||
from ..output import console, print_error, print_success, print_warning
|
||||
|
||||
config_path = Path(config_file or ctx.config_file)
|
||||
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
console.print(f"Validating [cyan]{config_path}[/cyan]...")
|
||||
|
||||
try:
|
||||
with open(config_path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
if data is None:
|
||||
print_error("Configuration file is empty")
|
||||
raise click.Abort()
|
||||
|
||||
from ...config import Config
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
|
||||
errors = []
|
||||
warnings = []
|
||||
|
||||
if not (1 <= config.server.port <= 65535):
|
||||
errors.append(f"Invalid server port: {config.server.port}")
|
||||
|
||||
valid_levels = ["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]
|
||||
if config.logging.level.upper() not in valid_levels:
|
||||
errors.append(f"Invalid logging level: {config.logging.level}")
|
||||
|
||||
if config.ssl.enabled:
|
||||
if not Path(config.ssl.cert_file).exists():
|
||||
warnings.append(f"SSL cert file not found: {config.ssl.cert_file}")
|
||||
if not Path(config.ssl.key_file).exists():
|
||||
warnings.append(f"SSL key file not found: {config.ssl.key_file}")
|
||||
|
||||
valid_extension_types = [
|
||||
"routing",
|
||||
"process_orchestration",
|
||||
"asgi_mount",
|
||||
]
|
||||
for ext in config.extensions:
|
||||
if ext.type not in valid_extension_types:
|
||||
warnings.append(f"Unknown extension type: {ext.type}")
|
||||
|
||||
if strict:
|
||||
known_top_level = {"http", "server", "ssl", "logging", "extensions"}
|
||||
for key in data.keys():
|
||||
if key not in known_top_level:
|
||||
warnings.append(f"Unknown top-level field: {key}")
|
||||
|
||||
if errors:
|
||||
for error in errors:
|
||||
print_error(error)
|
||||
raise click.Abort()
|
||||
|
||||
if warnings:
|
||||
for warning in warnings:
|
||||
print_warning(warning)
|
||||
|
||||
print_success("Configuration is valid!")
|
||||
|
||||
except yaml.YAMLError as e:
|
||||
print_error(f"YAML syntax error: {e}")
|
||||
raise click.Abort()
|
||||
except Exception as e:
|
||||
print_error(f"Validation error: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@config_cmd.command("show")
|
||||
@click.option(
|
||||
"-c",
|
||||
"--config",
|
||||
"config_file",
|
||||
default=None,
|
||||
help="Path to configuration file",
|
||||
)
|
||||
@click.option(
|
||||
"--format",
|
||||
"output_format",
|
||||
type=click.Choice(["yaml", "json", "table"]),
|
||||
default="yaml",
|
||||
help="Output format",
|
||||
)
|
||||
@click.option(
|
||||
"--section",
|
||||
"section",
|
||||
default=None,
|
||||
help="Show only a specific section (e.g., server, logging)",
|
||||
)
|
||||
@click.pass_obj
|
||||
def show_cmd(ctx: Any, config_file: Optional[str], output_format: str, section: Optional[str]) -> None:
|
||||
"""
|
||||
Display current configuration.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve config show
|
||||
pyserve config show --format json
|
||||
pyserve config show --section server
|
||||
"""
|
||||
from ..output import console, print_error
|
||||
|
||||
config_path = Path(config_file or ctx.config_file)
|
||||
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
try:
|
||||
with open(config_path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
if section:
|
||||
if section in data:
|
||||
data = {section: data[section]}
|
||||
else:
|
||||
print_error(f"Section '{section}' not found in configuration")
|
||||
raise click.Abort()
|
||||
|
||||
if output_format == "yaml":
|
||||
from rich.syntax import Syntax
|
||||
|
||||
yaml_str = yaml.dump(data, default_flow_style=False, sort_keys=False)
|
||||
syntax = Syntax(yaml_str, "yaml", theme="monokai", line_numbers=False)
|
||||
console.print(syntax)
|
||||
|
||||
elif output_format == "json":
|
||||
from rich.syntax import Syntax
|
||||
|
||||
json_str = json.dumps(data, indent=2)
|
||||
syntax = Syntax(json_str, "json", theme="monokai", line_numbers=False)
|
||||
console.print(syntax)
|
||||
|
||||
elif output_format == "table":
|
||||
from rich.tree import Tree
|
||||
|
||||
def build_tree(data: Any, tree: Any) -> None:
|
||||
if isinstance(data, dict):
|
||||
for key, value in data.items():
|
||||
if isinstance(value, (dict, list)):
|
||||
branch = tree.add(f"[cyan]{key}[/cyan]")
|
||||
build_tree(value, branch)
|
||||
else:
|
||||
tree.add(f"[cyan]{key}[/cyan]: [green]{value}[/green]")
|
||||
elif isinstance(data, list):
|
||||
for i, item in enumerate(data):
|
||||
if isinstance(item, (dict, list)):
|
||||
branch = tree.add(f"[dim][{i}][/dim]")
|
||||
build_tree(item, branch)
|
||||
else:
|
||||
tree.add(f"[dim][{i}][/dim] [green]{item}[/green]")
|
||||
|
||||
tree = Tree(f"[bold]Configuration: {config_path}[/bold]")
|
||||
build_tree(data, tree)
|
||||
console.print(tree)
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error reading configuration: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@config_cmd.command("get")
|
||||
@click.argument("key")
|
||||
@click.option(
|
||||
"-c",
|
||||
"--config",
|
||||
"config_file",
|
||||
default=None,
|
||||
help="Path to configuration file",
|
||||
)
|
||||
@click.pass_obj
|
||||
def get_cmd(ctx: Any, key: str, config_file: Optional[str]) -> None:
|
||||
"""
|
||||
Get a specific configuration value.
|
||||
|
||||
Use dot notation to access nested values.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve config get server.port
|
||||
pyserve config get logging.level
|
||||
pyserve config get extensions.0.type
|
||||
"""
|
||||
from ..output import console, print_error
|
||||
|
||||
config_path = Path(config_file or ctx.config_file)
|
||||
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
try:
|
||||
with open(config_path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
value = data
|
||||
for part in key.split("."):
|
||||
if isinstance(value, dict):
|
||||
if part in value:
|
||||
value = value[part]
|
||||
else:
|
||||
print_error(f"Key '{key}' not found")
|
||||
raise click.Abort()
|
||||
elif isinstance(value, list):
|
||||
try:
|
||||
index = int(part)
|
||||
value = value[index]
|
||||
except (ValueError, IndexError):
|
||||
print_error(f"Invalid index '{part}' in key '{key}'")
|
||||
raise click.Abort()
|
||||
else:
|
||||
print_error(f"Cannot access '{part}' in {type(value).__name__}")
|
||||
raise click.Abort()
|
||||
|
||||
if isinstance(value, (dict, list)):
|
||||
console.print(yaml.dump(value, default_flow_style=False))
|
||||
else:
|
||||
console.print(str(value))
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@config_cmd.command("set")
|
||||
@click.argument("key")
|
||||
@click.argument("value")
|
||||
@click.option(
|
||||
"-c",
|
||||
"--config",
|
||||
"config_file",
|
||||
default=None,
|
||||
help="Path to configuration file",
|
||||
)
|
||||
@click.pass_obj
|
||||
def set_cmd(ctx: Any, key: str, value: str, config_file: Optional[str]) -> None:
|
||||
"""
|
||||
Set a configuration value.
|
||||
|
||||
Use dot notation to access nested values.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve config set server.port 8080
|
||||
pyserve config set logging.level DEBUG
|
||||
"""
|
||||
from ..output import print_error, print_success
|
||||
|
||||
config_path = Path(config_file or ctx.config_file)
|
||||
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
try:
|
||||
with open(config_path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
parsed_value: Any
|
||||
if value.lower() == "true":
|
||||
parsed_value = True
|
||||
elif value.lower() == "false":
|
||||
parsed_value = False
|
||||
elif value.isdigit():
|
||||
parsed_value = int(value)
|
||||
else:
|
||||
try:
|
||||
parsed_value = float(value)
|
||||
except ValueError:
|
||||
parsed_value = value
|
||||
|
||||
parts = key.split(".")
|
||||
current = data
|
||||
for part in parts[:-1]:
|
||||
if isinstance(current, dict):
|
||||
if part not in current:
|
||||
current[part] = {}
|
||||
current = current[part]
|
||||
elif isinstance(current, list):
|
||||
index = int(part)
|
||||
current = current[index]
|
||||
|
||||
final_key = parts[-1]
|
||||
if isinstance(current, dict):
|
||||
current[final_key] = parsed_value
|
||||
elif isinstance(current, list):
|
||||
current[int(final_key)] = parsed_value
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
yaml.dump(data, f, default_flow_style=False, sort_keys=False)
|
||||
|
||||
print_success(f"Set {key} = {parsed_value}")
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@config_cmd.command("diff")
|
||||
@click.argument("file1", type=click.Path(exists=True))
|
||||
@click.argument("file2", type=click.Path(exists=True))
|
||||
def diff_cmd(file1: str, file2: str) -> None:
|
||||
"""
|
||||
Compare two configuration files.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve config diff config.yaml production.yaml
|
||||
"""
|
||||
from ..output import console, print_error
|
||||
|
||||
try:
|
||||
with open(file1) as f:
|
||||
data1 = yaml.safe_load(f)
|
||||
with open(file2) as f:
|
||||
data2 = yaml.safe_load(f)
|
||||
|
||||
def compare_dicts(d1: Any, d2: Any, path: str = "") -> list[tuple[str, str, Any, Any]]:
|
||||
differences: list[tuple[str, str, Any, Any]] = []
|
||||
|
||||
all_keys = set(d1.keys() if d1 else []) | set(d2.keys() if d2 else [])
|
||||
|
||||
for key in sorted(all_keys):
|
||||
current_path = f"{path}.{key}" if path else key
|
||||
v1 = d1.get(key) if d1 else None
|
||||
v2 = d2.get(key) if d2 else None
|
||||
|
||||
if key not in (d1 or {}):
|
||||
differences.append(("added", current_path, None, v2))
|
||||
elif key not in (d2 or {}):
|
||||
differences.append(("removed", current_path, v1, None))
|
||||
elif isinstance(v1, dict) and isinstance(v2, dict):
|
||||
differences.extend(compare_dicts(v1, v2, current_path))
|
||||
elif v1 != v2:
|
||||
differences.append(("changed", current_path, v1, v2))
|
||||
|
||||
return differences
|
||||
|
||||
differences = compare_dicts(data1, data2)
|
||||
|
||||
if not differences:
|
||||
console.print("[green]Files are identical[/green]")
|
||||
return
|
||||
|
||||
console.print(f"\n[bold]Differences between {file1} and {file2}:[/bold]\n")
|
||||
|
||||
for diff_type, path, v1, v2 in differences:
|
||||
if diff_type == "added":
|
||||
console.print(f" [green]+ {path}: {v2}[/green]")
|
||||
elif diff_type == "removed":
|
||||
console.print(f" [red]- {path}: {v1}[/red]")
|
||||
elif diff_type == "changed":
|
||||
console.print(f" [yellow]~ {path}:[/yellow]")
|
||||
console.print(f" [red]- {v1}[/red]")
|
||||
console.print(f" [green]+ {v2}[/green]")
|
||||
|
||||
console.print()
|
||||
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
@@ -0,0 +1,123 @@
|
||||
"""
|
||||
pyserve down - Stop all services
|
||||
"""
|
||||
|
||||
import signal
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("down")
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=30,
|
||||
type=int,
|
||||
help="Timeout in seconds for graceful shutdown",
|
||||
)
|
||||
@click.option(
|
||||
"-v",
|
||||
"--volumes",
|
||||
is_flag=True,
|
||||
help="Remove volumes/data",
|
||||
)
|
||||
@click.option(
|
||||
"--remove-orphans",
|
||||
is_flag=True,
|
||||
help="Remove orphaned services",
|
||||
)
|
||||
@click.pass_obj
|
||||
def down_cmd(
|
||||
ctx: Any,
|
||||
timeout: int,
|
||||
volumes: bool,
|
||||
remove_orphans: bool,
|
||||
) -> None:
|
||||
"""
|
||||
Stop and remove all services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve down # Stop all services
|
||||
pyserve down --timeout 60 # Extended shutdown timeout
|
||||
pyserve down -v # Remove volumes too
|
||||
"""
|
||||
from ..output import console, print_error, print_info, print_success, print_warning
|
||||
from ..state import StateManager
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
|
||||
if state_manager.is_daemon_running():
|
||||
daemon_pid = state_manager.get_daemon_pid()
|
||||
console.print(f"[bold]Stopping PyServe daemon (PID: {daemon_pid})...[/bold]")
|
||||
|
||||
try:
|
||||
import os
|
||||
|
||||
# FIXME: Please fix the cast usage here
|
||||
os.kill(cast(int, daemon_pid), signal.SIGTERM)
|
||||
|
||||
start_time = time.time()
|
||||
while time.time() - start_time < timeout:
|
||||
try:
|
||||
# FIXME: Please fix the cast usage here
|
||||
os.kill(cast(int, daemon_pid), 0)
|
||||
time.sleep(0.5)
|
||||
except ProcessLookupError:
|
||||
break
|
||||
else:
|
||||
print_warning("Graceful shutdown timed out, forcing...")
|
||||
try:
|
||||
# FIXME: Please fix the cast usage here
|
||||
os.kill(cast(int, daemon_pid), signal.SIGKILL)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
state_manager.clear_daemon_pid()
|
||||
print_success("PyServe daemon stopped")
|
||||
|
||||
except ProcessLookupError:
|
||||
print_info("Daemon was not running")
|
||||
state_manager.clear_daemon_pid()
|
||||
except PermissionError:
|
||||
print_error("Permission denied to stop daemon")
|
||||
raise click.Abort()
|
||||
else:
|
||||
services = state_manager.get_all_services()
|
||||
|
||||
if not services:
|
||||
print_info("No services are running")
|
||||
return
|
||||
|
||||
console.print("[bold]Stopping services...[/bold]")
|
||||
|
||||
from ...config import Config
|
||||
from .._runner import ServiceRunner
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
if config_path.exists():
|
||||
config = Config.from_yaml(str(config_path))
|
||||
else:
|
||||
config = Config()
|
||||
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
import asyncio
|
||||
|
||||
try:
|
||||
asyncio.run(runner.stop_all(timeout=timeout))
|
||||
print_success("All services stopped")
|
||||
except Exception as e:
|
||||
print_error(f"Error stopping services: {e}")
|
||||
|
||||
if volumes:
|
||||
console.print("Cleaning up state...")
|
||||
state_manager.clear()
|
||||
print_info("State cleared")
|
||||
|
||||
if remove_orphans:
|
||||
# This would remove services that are in state but not in config
|
||||
pass
|
||||
@@ -0,0 +1,161 @@
|
||||
"""
|
||||
pyserve health - Check health of services
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("health")
|
||||
@click.argument("services", nargs=-1)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=5,
|
||||
type=int,
|
||||
help="Health check timeout in seconds",
|
||||
)
|
||||
@click.option(
|
||||
"--format",
|
||||
"output_format",
|
||||
type=click.Choice(["table", "json"]),
|
||||
default="table",
|
||||
help="Output format",
|
||||
)
|
||||
@click.pass_obj
|
||||
def health_cmd(ctx: Any, services: tuple[str, ...], timeout: int, output_format: str) -> None:
|
||||
"""
|
||||
Check health of services.
|
||||
|
||||
Performs active health checks on running services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve health # Check all services
|
||||
pyserve health api admin # Check specific services
|
||||
pyserve health --format json # JSON output
|
||||
"""
|
||||
from ..output import console, print_error, print_info
|
||||
from ..state import StateManager
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
all_services = state_manager.get_all_services()
|
||||
|
||||
if services:
|
||||
all_services = {k: v for k, v in all_services.items() if k in services}
|
||||
|
||||
if not all_services:
|
||||
print_info("No services to check")
|
||||
return
|
||||
|
||||
results = asyncio.run(_check_health(all_services, timeout))
|
||||
|
||||
if output_format == "json":
|
||||
import json
|
||||
|
||||
console.print(json.dumps(results, indent=2))
|
||||
return
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ..output import format_health
|
||||
|
||||
table = Table(show_header=True, header_style="bold")
|
||||
table.add_column("SERVICE", style="cyan")
|
||||
table.add_column("HEALTH")
|
||||
table.add_column("CHECKS", justify="right")
|
||||
table.add_column("LAST CHECK", style="dim")
|
||||
table.add_column("RESPONSE TIME", justify="right")
|
||||
|
||||
for name, result in results.items():
|
||||
health_str = format_health(result["status"])
|
||||
checks = f"{result['successes']}/{result['total']}"
|
||||
last_check = result.get("last_check", "-")
|
||||
response_time = f"{result['response_time_ms']:.0f}ms" if result.get("response_time_ms") else "-"
|
||||
|
||||
table.add_row(name, health_str, checks, last_check, response_time)
|
||||
|
||||
console.print()
|
||||
console.print(table)
|
||||
console.print()
|
||||
|
||||
healthy = sum(1 for r in results.values() if r["status"] == "healthy")
|
||||
unhealthy = sum(1 for r in results.values() if r["status"] == "unhealthy")
|
||||
|
||||
if unhealthy:
|
||||
print_error(f"{unhealthy} service(s) unhealthy")
|
||||
raise SystemExit(1)
|
||||
else:
|
||||
from ..output import print_success
|
||||
|
||||
print_success(f"All {healthy} service(s) healthy")
|
||||
|
||||
|
||||
async def _check_health(services: dict, timeout: int) -> dict:
|
||||
import time
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
return {name: {"status": "unknown", "error": "httpx not installed"} for name in services}
|
||||
|
||||
results = {}
|
||||
|
||||
for name, service in services.items():
|
||||
if service.state != "running" or not service.port:
|
||||
results[name] = {
|
||||
"status": "unknown",
|
||||
"successes": 0,
|
||||
"total": 0,
|
||||
"error": "Service not running",
|
||||
}
|
||||
continue
|
||||
|
||||
health_path = "/health"
|
||||
url = f"http://127.0.0.1:{service.port}{health_path}"
|
||||
|
||||
start_time = time.time()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
resp = await client.get(url)
|
||||
response_time = (time.time() - start_time) * 1000
|
||||
|
||||
if resp.status_code < 500:
|
||||
results[name] = {
|
||||
"status": "healthy",
|
||||
"successes": 1,
|
||||
"total": 1,
|
||||
"response_time_ms": response_time,
|
||||
"last_check": "just now",
|
||||
"status_code": resp.status_code,
|
||||
}
|
||||
else:
|
||||
results[name] = {
|
||||
"status": "unhealthy",
|
||||
"successes": 0,
|
||||
"total": 1,
|
||||
"response_time_ms": response_time,
|
||||
"last_check": "just now",
|
||||
"status_code": resp.status_code,
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
results[name] = {
|
||||
"status": "unhealthy",
|
||||
"successes": 0,
|
||||
"total": 1,
|
||||
"error": "timeout",
|
||||
"last_check": "just now",
|
||||
}
|
||||
except Exception as e:
|
||||
results[name] = {
|
||||
"status": "unhealthy",
|
||||
"successes": 0,
|
||||
"total": 1,
|
||||
"error": str(e),
|
||||
"last_check": "just now",
|
||||
}
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,432 @@
|
||||
"""
|
||||
pyserve init - Initialize a new pyserve project
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
|
||||
TEMPLATES = {
|
||||
"basic": {
|
||||
"description": "Basic configuration with static files and routing",
|
||||
"filename": "config.yaml",
|
||||
},
|
||||
"orchestration": {
|
||||
"description": "Process orchestration with multiple ASGI/WSGI apps",
|
||||
"filename": "config.yaml",
|
||||
},
|
||||
"asgi": {
|
||||
"description": "ASGI mount configuration for in-process apps",
|
||||
"filename": "config.yaml",
|
||||
},
|
||||
"full": {
|
||||
"description": "Full configuration with all features",
|
||||
"filename": "config.yaml",
|
||||
},
|
||||
}
|
||||
|
||||
BASIC_TEMPLATE = """\
|
||||
# PyServe Configuration
|
||||
# Generated by: pyserve init
|
||||
|
||||
http:
|
||||
static_dir: ./static
|
||||
templates_dir: ./templates
|
||||
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
backlog: 100
|
||||
proxy_timeout: 30.0
|
||||
|
||||
ssl:
|
||||
enabled: false
|
||||
cert_file: ./ssl/cert.pem
|
||||
key_file: ./ssl/key.pem
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
show_module: true
|
||||
timestamp_format: "%Y-%m-%d %H:%M:%S"
|
||||
console:
|
||||
level: INFO
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
files:
|
||||
- path: ./logs/pyserve.log
|
||||
level: INFO
|
||||
format:
|
||||
type: standard
|
||||
use_colors: false
|
||||
|
||||
extensions:
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
# Health check endpoint
|
||||
"=/health":
|
||||
return: "200 OK"
|
||||
content_type: "text/plain"
|
||||
|
||||
# Static files
|
||||
"^/static/":
|
||||
root: "./static"
|
||||
strip_prefix: "/static"
|
||||
|
||||
# Default fallback
|
||||
"__default__":
|
||||
spa_fallback: true
|
||||
root: "./static"
|
||||
index_file: "index.html"
|
||||
"""
|
||||
|
||||
ORCHESTRATION_TEMPLATE = """\
|
||||
# PyServe Process Orchestration Configuration
|
||||
# Generated by: pyserve init --template orchestration
|
||||
#
|
||||
# This configuration runs multiple ASGI/WSGI apps as isolated processes
|
||||
# with automatic health monitoring and restart.
|
||||
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
backlog: 2048
|
||||
proxy_timeout: 60.0
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
files:
|
||||
- path: ./logs/pyserve.log
|
||||
level: DEBUG
|
||||
format:
|
||||
type: standard
|
||||
use_colors: false
|
||||
|
||||
extensions:
|
||||
# Process Orchestration - runs each app in its own process
|
||||
- type: process_orchestration
|
||||
config:
|
||||
port_range: [9000, 9999]
|
||||
health_check_enabled: true
|
||||
proxy_timeout: 60.0
|
||||
|
||||
apps:
|
||||
# Example: FastAPI application
|
||||
- name: api
|
||||
path: /api
|
||||
app_path: myapp.api:app
|
||||
module_path: "."
|
||||
workers: 2
|
||||
health_check_path: /health
|
||||
health_check_interval: 10.0
|
||||
health_check_timeout: 5.0
|
||||
health_check_retries: 3
|
||||
max_restart_count: 5
|
||||
restart_delay: 1.0
|
||||
strip_path: true
|
||||
env:
|
||||
APP_ENV: "production"
|
||||
|
||||
# Example: Flask application (WSGI)
|
||||
# - name: admin
|
||||
# path: /admin
|
||||
# app_path: myapp.admin:app
|
||||
# app_type: wsgi
|
||||
# module_path: "."
|
||||
# workers: 1
|
||||
# health_check_path: /health
|
||||
# strip_path: true
|
||||
|
||||
# Static files routing
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
"=/health":
|
||||
return: "200 OK"
|
||||
content_type: "text/plain"
|
||||
|
||||
"^/static/":
|
||||
root: "./static"
|
||||
strip_prefix: "/static"
|
||||
"""
|
||||
|
||||
ASGI_TEMPLATE = """\
|
||||
# PyServe ASGI Mount Configuration
|
||||
# Generated by: pyserve init --template asgi
|
||||
#
|
||||
# This configuration mounts ASGI apps in-process (like ASGI Lifespan).
|
||||
# More efficient but apps share the same process.
|
||||
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
backlog: 100
|
||||
proxy_timeout: 30.0
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
files:
|
||||
- path: ./logs/pyserve.log
|
||||
level: DEBUG
|
||||
|
||||
extensions:
|
||||
- type: asgi_mount
|
||||
config:
|
||||
mounts:
|
||||
# FastAPI app mounted at /api
|
||||
- path: /api
|
||||
app: myapp.api:app
|
||||
# factory: false # Set to true if app is a factory function
|
||||
|
||||
# Starlette app mounted at /web
|
||||
# - path: /web
|
||||
# app: myapp.web:app
|
||||
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
"=/health":
|
||||
return: "200 OK"
|
||||
content_type: "text/plain"
|
||||
|
||||
"^/static/":
|
||||
root: "./static"
|
||||
strip_prefix: "/static"
|
||||
|
||||
"__default__":
|
||||
spa_fallback: true
|
||||
root: "./static"
|
||||
index_file: "index.html"
|
||||
"""
|
||||
|
||||
FULL_TEMPLATE = """\
|
||||
# PyServe Full Configuration
|
||||
# Generated by: pyserve init --template full
|
||||
#
|
||||
# Comprehensive configuration showcasing all PyServe features.
|
||||
|
||||
http:
|
||||
static_dir: ./static
|
||||
templates_dir: ./templates
|
||||
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
backlog: 2048
|
||||
default_root: false
|
||||
proxy_timeout: 60.0
|
||||
redirect_instructions:
|
||||
"/old-path": "/new-path"
|
||||
|
||||
ssl:
|
||||
enabled: false
|
||||
cert_file: ./ssl/cert.pem
|
||||
key_file: ./ssl/key.pem
|
||||
|
||||
logging:
|
||||
level: INFO
|
||||
console_output: true
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
show_module: true
|
||||
timestamp_format: "%Y-%m-%d %H:%M:%S"
|
||||
console:
|
||||
level: DEBUG
|
||||
format:
|
||||
type: standard
|
||||
use_colors: true
|
||||
files:
|
||||
# Main log file
|
||||
- path: ./logs/pyserve.log
|
||||
level: DEBUG
|
||||
format:
|
||||
type: standard
|
||||
use_colors: false
|
||||
|
||||
# JSON logs for log aggregation
|
||||
- path: ./logs/pyserve.json
|
||||
level: INFO
|
||||
format:
|
||||
type: json
|
||||
|
||||
# Access logs
|
||||
- path: ./logs/access.log
|
||||
level: INFO
|
||||
loggers: ["pyserve.access"]
|
||||
max_bytes: 10485760 # 10MB
|
||||
backup_count: 10
|
||||
|
||||
extensions:
|
||||
# Process Orchestration for background services
|
||||
- type: process_orchestration
|
||||
config:
|
||||
port_range: [9000, 9999]
|
||||
health_check_enabled: true
|
||||
proxy_timeout: 60.0
|
||||
|
||||
apps:
|
||||
- name: api
|
||||
path: /api
|
||||
app_path: myapp.api:app
|
||||
module_path: "."
|
||||
workers: 2
|
||||
health_check_path: /health
|
||||
strip_path: true
|
||||
env:
|
||||
APP_ENV: "production"
|
||||
|
||||
# Advanced routing with regex
|
||||
- type: routing
|
||||
config:
|
||||
regex_locations:
|
||||
# API versioning
|
||||
"~^/api/v(?P<version>\\\\d+)/":
|
||||
proxy_pass: "http://localhost:9001"
|
||||
headers:
|
||||
- "API-Version: {version}"
|
||||
- "X-Forwarded-For: $remote_addr"
|
||||
|
||||
# Static files with caching
|
||||
"~*\\\\.(js|css|png|jpg|gif|ico|svg|woff2?)$":
|
||||
root: "./static"
|
||||
cache_control: "public, max-age=31536000"
|
||||
headers:
|
||||
- "Access-Control-Allow-Origin: *"
|
||||
|
||||
# Health check
|
||||
"=/health":
|
||||
return: "200 OK"
|
||||
content_type: "text/plain"
|
||||
|
||||
# Static files
|
||||
"^/static/":
|
||||
root: "./static"
|
||||
strip_prefix: "/static"
|
||||
|
||||
# SPA fallback
|
||||
"__default__":
|
||||
spa_fallback: true
|
||||
root: "./static"
|
||||
index_file: "index.html"
|
||||
"""
|
||||
|
||||
|
||||
def get_template_content(template: str) -> str:
|
||||
templates = {
|
||||
"basic": BASIC_TEMPLATE,
|
||||
"orchestration": ORCHESTRATION_TEMPLATE,
|
||||
"asgi": ASGI_TEMPLATE,
|
||||
"full": FULL_TEMPLATE,
|
||||
}
|
||||
return templates.get(template, BASIC_TEMPLATE)
|
||||
|
||||
|
||||
@click.command("init")
|
||||
@click.option(
|
||||
"-t",
|
||||
"--template",
|
||||
"template",
|
||||
type=click.Choice(list(TEMPLATES.keys())),
|
||||
default="basic",
|
||||
help="Configuration template to use",
|
||||
)
|
||||
@click.option(
|
||||
"-o",
|
||||
"--output",
|
||||
"output_file",
|
||||
default="config.yaml",
|
||||
help="Output file path (default: config.yaml)",
|
||||
)
|
||||
@click.option(
|
||||
"-f",
|
||||
"--force",
|
||||
is_flag=True,
|
||||
help="Overwrite existing configuration",
|
||||
)
|
||||
@click.option(
|
||||
"--list-templates",
|
||||
is_flag=True,
|
||||
help="List available templates",
|
||||
)
|
||||
@click.pass_context
|
||||
def init_cmd(
|
||||
ctx: click.Context,
|
||||
template: str,
|
||||
output_file: str,
|
||||
force: bool,
|
||||
list_templates: bool,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize a new pyserve project.
|
||||
|
||||
Creates a configuration file with sensible defaults and directory structure.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve init # Basic configuration
|
||||
pyserve init -t orchestration # Process orchestration setup
|
||||
pyserve init -t asgi # ASGI mount setup
|
||||
pyserve init -t full # All features
|
||||
pyserve init -o production.yaml # Custom output file
|
||||
"""
|
||||
from ..output import console, print_info, print_success, print_warning
|
||||
|
||||
if list_templates:
|
||||
console.print("\n[bold]Available Templates:[/bold]\n")
|
||||
for name, info in TEMPLATES.items():
|
||||
console.print(f" [cyan]{name:15}[/cyan] - {info['description']}")
|
||||
console.print()
|
||||
return
|
||||
|
||||
output_path = Path(output_file)
|
||||
|
||||
if output_path.exists() and not force:
|
||||
print_warning(f"Configuration file '{output_file}' already exists.")
|
||||
if not click.confirm("Do you want to overwrite it?"):
|
||||
raise click.Abort()
|
||||
|
||||
dirs_to_create = ["static", "templates", "logs"]
|
||||
if template == "orchestration":
|
||||
dirs_to_create.append("apps")
|
||||
|
||||
for dir_name in dirs_to_create:
|
||||
dir_path = Path(dir_name)
|
||||
if not dir_path.exists():
|
||||
dir_path.mkdir(parents=True)
|
||||
print_info(f"Created directory: {dir_name}/")
|
||||
|
||||
state_dir = Path(".pyserve")
|
||||
if not state_dir.exists():
|
||||
state_dir.mkdir()
|
||||
print_info("Created directory: .pyserve/")
|
||||
|
||||
content = get_template_content(template)
|
||||
output_path.write_text(content)
|
||||
|
||||
print_success(f"Created configuration file: {output_file}")
|
||||
print_info(f"Template: {template}")
|
||||
|
||||
gitignore_path = Path(".pyserve/.gitignore")
|
||||
if not gitignore_path.exists():
|
||||
gitignore_path.write_text("*\n!.gitignore\n")
|
||||
|
||||
console.print()
|
||||
console.print("[bold]Next steps:[/bold]")
|
||||
console.print(f" 1. Edit [cyan]{output_file}[/cyan] to configure your services")
|
||||
console.print(" 2. Run [cyan]pyserve config validate[/cyan] to check configuration")
|
||||
console.print(" 3. Run [cyan]pyserve up[/cyan] to start services")
|
||||
console.print()
|
||||
@@ -0,0 +1,280 @@
|
||||
"""
|
||||
pyserve logs - View service logs
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("logs")
|
||||
@click.argument("services", nargs=-1)
|
||||
@click.option(
|
||||
"-f",
|
||||
"--follow",
|
||||
is_flag=True,
|
||||
help="Follow log output",
|
||||
)
|
||||
@click.option(
|
||||
"--tail",
|
||||
"tail",
|
||||
default=100,
|
||||
type=int,
|
||||
help="Number of lines to show from the end",
|
||||
)
|
||||
@click.option(
|
||||
"--since",
|
||||
"since",
|
||||
default=None,
|
||||
help="Show logs since timestamp (e.g., '10m', '1h', '2024-01-01')",
|
||||
)
|
||||
@click.option(
|
||||
"--until",
|
||||
"until_time",
|
||||
default=None,
|
||||
help="Show logs until timestamp",
|
||||
)
|
||||
@click.option(
|
||||
"-t",
|
||||
"--timestamps",
|
||||
is_flag=True,
|
||||
help="Show timestamps",
|
||||
)
|
||||
@click.option(
|
||||
"--no-color",
|
||||
is_flag=True,
|
||||
help="Disable colored output",
|
||||
)
|
||||
@click.option(
|
||||
"--filter",
|
||||
"filter_pattern",
|
||||
default=None,
|
||||
help="Filter logs by pattern",
|
||||
)
|
||||
@click.pass_obj
|
||||
def logs_cmd(
|
||||
ctx: Any,
|
||||
services: tuple[str, ...],
|
||||
follow: bool,
|
||||
tail: int,
|
||||
since: Optional[str],
|
||||
until_time: Optional[str],
|
||||
timestamps: bool,
|
||||
no_color: bool,
|
||||
filter_pattern: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
View service logs.
|
||||
|
||||
If no services are specified, shows logs from all services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve logs # All logs
|
||||
pyserve logs api # Logs from api service
|
||||
pyserve logs api admin # Logs from multiple services
|
||||
pyserve logs -f # Follow logs
|
||||
pyserve logs --tail 50 # Last 50 lines
|
||||
pyserve logs --since "10m" # Logs from last 10 minutes
|
||||
"""
|
||||
from ..output import print_info
|
||||
from ..state import StateManager
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
|
||||
if services:
|
||||
log_files = [(name, state_manager.get_service_log_file(name)) for name in services]
|
||||
else:
|
||||
all_services = state_manager.get_all_services()
|
||||
if not all_services:
|
||||
main_log = Path("logs/pyserve.log")
|
||||
if main_log.exists():
|
||||
log_files = [("pyserve", main_log)]
|
||||
else:
|
||||
print_info("No logs available. Start services with 'pyserve up'")
|
||||
return
|
||||
else:
|
||||
log_files = [(name, state_manager.get_service_log_file(name)) for name in all_services]
|
||||
|
||||
existing_logs = [(name, path) for name, path in log_files if path.exists()]
|
||||
|
||||
if not existing_logs:
|
||||
print_info("No log files found")
|
||||
return
|
||||
|
||||
since_time = _parse_time(since) if since else None
|
||||
until_timestamp = _parse_time(until_time) if until_time else None
|
||||
|
||||
colors = ["cyan", "green", "yellow", "blue", "magenta"]
|
||||
service_colors = {name: colors[i % len(colors)] for i, (name, _) in enumerate(existing_logs)}
|
||||
|
||||
if follow:
|
||||
asyncio.run(
|
||||
_follow_logs(
|
||||
existing_logs,
|
||||
service_colors,
|
||||
timestamps,
|
||||
no_color,
|
||||
filter_pattern,
|
||||
)
|
||||
)
|
||||
else:
|
||||
_read_logs(
|
||||
existing_logs,
|
||||
service_colors,
|
||||
tail,
|
||||
since_time,
|
||||
until_timestamp,
|
||||
timestamps,
|
||||
no_color,
|
||||
filter_pattern,
|
||||
)
|
||||
|
||||
|
||||
def _parse_time(time_str: str) -> Optional[float]:
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
# Relative time (e.g., "10m", "1h", "2d")
|
||||
match = re.match(r"^(\d+)([smhd])$", time_str)
|
||||
if match:
|
||||
value = int(match.group(1))
|
||||
unit = match.group(2)
|
||||
units = {"s": 1, "m": 60, "h": 3600, "d": 86400}
|
||||
return time.time() - (value * units[unit])
|
||||
|
||||
# Relative phrase (e.g., "10m ago")
|
||||
match = re.match(r"^(\d+)([smhd])\s+ago$", time_str)
|
||||
if match:
|
||||
value = int(match.group(1))
|
||||
unit = match.group(2)
|
||||
units = {"s": 1, "m": 60, "h": 3600, "d": 86400}
|
||||
return time.time() - (value * units[unit])
|
||||
|
||||
# ISO format
|
||||
try:
|
||||
dt = datetime.fromisoformat(time_str)
|
||||
return dt.timestamp()
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _read_logs(
|
||||
log_files: list[tuple[str, Path]],
|
||||
service_colors: dict[str, str],
|
||||
tail: int,
|
||||
since_time: Optional[float],
|
||||
until_time: Optional[float],
|
||||
timestamps: bool,
|
||||
no_color: bool,
|
||||
filter_pattern: Optional[str],
|
||||
) -> None:
|
||||
import re
|
||||
|
||||
from ..output import console
|
||||
|
||||
all_lines = []
|
||||
|
||||
for service_name, log_path in log_files:
|
||||
try:
|
||||
with open(log_path) as f:
|
||||
lines = f.readlines()
|
||||
|
||||
# Take last N lines
|
||||
lines = lines[-tail:] if tail else lines
|
||||
|
||||
for line in lines:
|
||||
line = line.rstrip()
|
||||
if not line:
|
||||
continue
|
||||
|
||||
if filter_pattern and filter_pattern not in line:
|
||||
continue
|
||||
|
||||
line_time = None
|
||||
timestamp_match = re.match(r"^(\d{4}-\d{2}-\d{2}[T ]\d{2}:\d{2}:\d{2})", line)
|
||||
if timestamp_match:
|
||||
try:
|
||||
from datetime import datetime
|
||||
|
||||
line_time = datetime.fromisoformat(timestamp_match.group(1).replace(" ", "T")).timestamp()
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if since_time and line_time and line_time < since_time:
|
||||
continue
|
||||
if until_time and line_time and line_time > until_time:
|
||||
continue
|
||||
|
||||
all_lines.append((line_time or 0, service_name, line))
|
||||
|
||||
except Exception as e:
|
||||
console.print(f"[red]Error reading {log_path}: {e}[/red]")
|
||||
|
||||
all_lines.sort(key=lambda x: x[0])
|
||||
|
||||
for _, service_name, line in all_lines:
|
||||
if len(log_files) > 1:
|
||||
# Multiple services - prefix with service name
|
||||
if no_color:
|
||||
console.print(f"{service_name} | {line}")
|
||||
else:
|
||||
color = service_colors.get(service_name, "white")
|
||||
console.print(f"[{color}]{service_name}[/{color}] | {line}")
|
||||
else:
|
||||
console.print(line)
|
||||
|
||||
|
||||
async def _follow_logs(
|
||||
log_files: list[tuple[str, Path]],
|
||||
service_colors: dict[str, str],
|
||||
timestamps: bool,
|
||||
no_color: bool,
|
||||
filter_pattern: Optional[str],
|
||||
) -> None:
|
||||
from ..output import console
|
||||
|
||||
positions = {}
|
||||
for service_name, log_path in log_files:
|
||||
if log_path.exists():
|
||||
positions[service_name] = log_path.stat().st_size
|
||||
else:
|
||||
positions[service_name] = 0
|
||||
|
||||
console.print("[dim]Following logs... Press Ctrl+C to stop[/dim]\n")
|
||||
|
||||
try:
|
||||
while True:
|
||||
for service_name, log_path in log_files:
|
||||
if not log_path.exists():
|
||||
continue
|
||||
|
||||
current_size = log_path.stat().st_size
|
||||
if current_size > positions[service_name]:
|
||||
with open(log_path) as f:
|
||||
f.seek(positions[service_name])
|
||||
new_content = f.read()
|
||||
positions[service_name] = f.tell()
|
||||
|
||||
for line in new_content.splitlines():
|
||||
if filter_pattern and filter_pattern not in line:
|
||||
continue
|
||||
|
||||
if len(log_files) > 1:
|
||||
if no_color:
|
||||
console.print(f"{service_name} | {line}")
|
||||
else:
|
||||
color = service_colors.get(service_name, "white")
|
||||
console.print(f"[{color}]{service_name}[/{color}] | {line}")
|
||||
else:
|
||||
console.print(line)
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n[dim]Stopped following logs[/dim]")
|
||||
@@ -0,0 +1,88 @@
|
||||
"""
|
||||
pyserve scale - Scale services
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("scale")
|
||||
@click.argument("scales", nargs=-1, required=True)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=60,
|
||||
type=int,
|
||||
help="Timeout in seconds for scaling operation",
|
||||
)
|
||||
@click.option(
|
||||
"--no-wait",
|
||||
is_flag=True,
|
||||
help="Don't wait for services to be ready",
|
||||
)
|
||||
@click.pass_obj
|
||||
def scale_cmd(ctx: Any, scales: tuple[str, ...], timeout: int, no_wait: bool) -> None:
|
||||
"""
|
||||
Scale services to specified number of workers.
|
||||
|
||||
Use SERVICE=NUM format to specify scaling.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve scale api=4 # Scale api to 4 workers
|
||||
pyserve scale api=4 admin=2 # Scale multiple services
|
||||
"""
|
||||
from ...config import Config
|
||||
from .._runner import ServiceRunner
|
||||
from ..output import console, print_error, print_info, print_success
|
||||
from ..state import StateManager
|
||||
|
||||
scale_map = {}
|
||||
for scale in scales:
|
||||
try:
|
||||
service, num = scale.split("=")
|
||||
scale_map[service] = int(num)
|
||||
except ValueError:
|
||||
print_error(f"Invalid scale format: {scale}. Use SERVICE=NUM")
|
||||
raise click.Abort()
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
|
||||
all_services = state_manager.get_all_services()
|
||||
for service in scale_map:
|
||||
if service not in all_services:
|
||||
print_error(f"Service '{service}' not found")
|
||||
raise click.Abort()
|
||||
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
console.print("[bold]Scaling services...[/bold]")
|
||||
|
||||
async def do_scale() -> None:
|
||||
for service, workers in scale_map.items():
|
||||
current = all_services[service].workers or 1
|
||||
print_info(f"Scaling {service}: {current} → {workers} workers")
|
||||
|
||||
try:
|
||||
success = await runner.scale_service(service, workers, timeout=timeout, wait=not no_wait)
|
||||
if success:
|
||||
print_success(f"Scaled {service} to {workers} workers")
|
||||
else:
|
||||
print_error(f"Failed to scale {service}")
|
||||
except Exception as e:
|
||||
print_error(f"Error scaling {service}: {e}")
|
||||
|
||||
try:
|
||||
asyncio.run(do_scale())
|
||||
except Exception as e:
|
||||
print_error(f"Scaling failed: {e}")
|
||||
raise click.Abort()
|
||||
@@ -0,0 +1,190 @@
|
||||
"""
|
||||
pyserve start/stop/restart - Service management commands
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("start")
|
||||
@click.argument("services", nargs=-1, required=True)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=60,
|
||||
type=int,
|
||||
help="Timeout in seconds for service startup",
|
||||
)
|
||||
@click.pass_obj
|
||||
def start_cmd(ctx: Any, services: tuple[str, ...], timeout: int) -> None:
|
||||
"""
|
||||
Start one or more services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve start api # Start api service
|
||||
pyserve start api admin # Start multiple services
|
||||
"""
|
||||
from ...config import Config
|
||||
from .._runner import ServiceRunner
|
||||
from ..output import console, print_error, print_success
|
||||
from ..state import StateManager
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
console.print(f"[bold]Starting services: {', '.join(services)}[/bold]")
|
||||
|
||||
async def do_start() -> Dict[str, bool]:
|
||||
results: Dict[str, bool] = {}
|
||||
for service in services:
|
||||
try:
|
||||
success = await runner.start_service(service, timeout=timeout)
|
||||
results[service] = success
|
||||
if success:
|
||||
print_success(f"Started {service}")
|
||||
else:
|
||||
print_error(f"Failed to start {service}")
|
||||
except Exception as e:
|
||||
print_error(f"Error starting {service}: {e}")
|
||||
results[service] = False
|
||||
return results
|
||||
|
||||
try:
|
||||
results = asyncio.run(do_start())
|
||||
if not all(results.values()):
|
||||
raise click.Abort()
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@click.command("stop")
|
||||
@click.argument("services", nargs=-1, required=True)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=30,
|
||||
type=int,
|
||||
help="Timeout in seconds for graceful shutdown",
|
||||
)
|
||||
@click.option(
|
||||
"-f",
|
||||
"--force",
|
||||
is_flag=True,
|
||||
help="Force stop (SIGKILL)",
|
||||
)
|
||||
@click.pass_obj
|
||||
def stop_cmd(ctx: Any, services: tuple[str, ...], timeout: int, force: bool) -> None:
|
||||
"""
|
||||
Stop one or more services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve stop api # Stop api service
|
||||
pyserve stop api admin # Stop multiple services
|
||||
pyserve stop api --force # Force stop
|
||||
"""
|
||||
from ...config import Config
|
||||
from .._runner import ServiceRunner
|
||||
from ..output import console, print_error, print_success
|
||||
from ..state import StateManager
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
config = Config.from_yaml(str(config_path)) if config_path.exists() else Config()
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
console.print(f"[bold]Stopping services: {', '.join(services)}[/bold]")
|
||||
|
||||
async def do_stop() -> Dict[str, bool]:
|
||||
results: Dict[str, bool] = {}
|
||||
for service in services:
|
||||
try:
|
||||
success = await runner.stop_service(service, timeout=timeout, force=force)
|
||||
results[service] = success
|
||||
if success:
|
||||
print_success(f"Stopped {service}")
|
||||
else:
|
||||
print_error(f"Failed to stop {service}")
|
||||
except Exception as e:
|
||||
print_error(f"Error stopping {service}: {e}")
|
||||
results[service] = False
|
||||
return results
|
||||
|
||||
try:
|
||||
results = asyncio.run(do_stop())
|
||||
if not all(results.values()):
|
||||
raise click.Abort()
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
|
||||
@click.command("restart")
|
||||
@click.argument("services", nargs=-1, required=True)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=60,
|
||||
type=int,
|
||||
help="Timeout in seconds for restart",
|
||||
)
|
||||
@click.pass_obj
|
||||
def restart_cmd(ctx: Any, services: tuple[str, ...], timeout: int) -> None:
|
||||
"""
|
||||
Restart one or more services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve restart api # Restart api service
|
||||
pyserve restart api admin # Restart multiple services
|
||||
"""
|
||||
from ...config import Config
|
||||
from .._runner import ServiceRunner
|
||||
from ..output import console, print_error, print_success
|
||||
from ..state import StateManager
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
raise click.Abort()
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
console.print(f"[bold]Restarting services: {', '.join(services)}[/bold]")
|
||||
|
||||
async def do_restart() -> Dict[str, bool]:
|
||||
results = {}
|
||||
for service in services:
|
||||
try:
|
||||
success = await runner.restart_service(service, timeout=timeout)
|
||||
results[service] = success
|
||||
if success:
|
||||
print_success(f"Restarted {service}")
|
||||
else:
|
||||
print_error(f"Failed to restart {service}")
|
||||
except Exception as e:
|
||||
print_error(f"Error restarting {service}: {e}")
|
||||
results[service] = False
|
||||
return results
|
||||
|
||||
try:
|
||||
results = asyncio.run(do_restart())
|
||||
if not all(results.values()):
|
||||
raise click.Abort()
|
||||
except Exception as e:
|
||||
print_error(f"Error: {e}")
|
||||
raise click.Abort()
|
||||
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
pyserve ps / status - Show service status
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("ps")
|
||||
@click.argument("services", nargs=-1)
|
||||
@click.option(
|
||||
"-a",
|
||||
"--all",
|
||||
"show_all",
|
||||
is_flag=True,
|
||||
help="Show all services (including stopped)",
|
||||
)
|
||||
@click.option(
|
||||
"-q",
|
||||
"--quiet",
|
||||
is_flag=True,
|
||||
help="Only show service names",
|
||||
)
|
||||
@click.option(
|
||||
"--format",
|
||||
"output_format",
|
||||
type=click.Choice(["table", "json", "yaml"]),
|
||||
default="table",
|
||||
help="Output format",
|
||||
)
|
||||
@click.option(
|
||||
"--filter",
|
||||
"filter_status",
|
||||
default=None,
|
||||
help="Filter by status (running, stopped, failed)",
|
||||
)
|
||||
@click.pass_obj
|
||||
def ps_cmd(
|
||||
ctx: Any,
|
||||
services: tuple[str, ...],
|
||||
show_all: bool,
|
||||
quiet: bool,
|
||||
output_format: str,
|
||||
filter_status: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
Show status of services.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve ps # Show running services
|
||||
pyserve ps -a # Show all services
|
||||
pyserve ps api admin # Show specific services
|
||||
pyserve ps --format json # JSON output
|
||||
pyserve ps --filter running # Filter by status
|
||||
"""
|
||||
from ..output import (
|
||||
console,
|
||||
create_services_table,
|
||||
format_health,
|
||||
format_status,
|
||||
format_uptime,
|
||||
print_info,
|
||||
)
|
||||
from ..state import StateManager
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
all_services = state_manager.get_all_services()
|
||||
|
||||
# Check if daemon is running
|
||||
daemon_running = state_manager.is_daemon_running()
|
||||
|
||||
# Filter services
|
||||
if services:
|
||||
all_services = {k: v for k, v in all_services.items() if k in services}
|
||||
|
||||
if filter_status:
|
||||
all_services = {k: v for k, v in all_services.items() if v.state.lower() == filter_status.lower()}
|
||||
|
||||
if not show_all:
|
||||
# By default, show only running/starting/failed services
|
||||
all_services = {k: v for k, v in all_services.items() if v.state.lower() in ("running", "starting", "stopping", "failed", "restarting")}
|
||||
|
||||
if not all_services:
|
||||
if daemon_running:
|
||||
print_info("No services found. Daemon is running but no services are configured.")
|
||||
else:
|
||||
print_info("No services running. Use 'pyserve up' to start services.")
|
||||
return
|
||||
|
||||
if quiet:
|
||||
for name in all_services:
|
||||
click.echo(name)
|
||||
return
|
||||
|
||||
if output_format == "json":
|
||||
data = {name: svc.to_dict() for name, svc in all_services.items()}
|
||||
console.print(json.dumps(data, indent=2))
|
||||
return
|
||||
|
||||
if output_format == "yaml":
|
||||
import yaml
|
||||
|
||||
data = {name: svc.to_dict() for name, svc in all_services.items()}
|
||||
console.print(yaml.dump(data, default_flow_style=False))
|
||||
return
|
||||
|
||||
table = create_services_table()
|
||||
|
||||
for name, service in sorted(all_services.items()):
|
||||
ports = f"{service.port}" if service.port else "-"
|
||||
uptime = format_uptime(service.uptime) if service.state == "running" else "-"
|
||||
health = format_health(service.health.status if service.state == "running" else "-")
|
||||
pid = str(service.pid) if service.pid else "-"
|
||||
workers = f"{service.workers}" if service.workers else "-"
|
||||
|
||||
table.add_row(
|
||||
name,
|
||||
format_status(service.state),
|
||||
ports,
|
||||
uptime,
|
||||
health,
|
||||
pid,
|
||||
workers,
|
||||
)
|
||||
|
||||
console.print()
|
||||
console.print(table)
|
||||
console.print()
|
||||
|
||||
total = len(all_services)
|
||||
running = sum(1 for s in all_services.values() if s.state == "running")
|
||||
failed = sum(1 for s in all_services.values() if s.state == "failed")
|
||||
|
||||
summary_parts = [f"[bold]{total}[/bold] service(s)"]
|
||||
if running:
|
||||
summary_parts.append(f"[green]{running} running[/green]")
|
||||
if failed:
|
||||
summary_parts.append(f"[red]{failed} failed[/red]")
|
||||
if total - running - failed > 0:
|
||||
summary_parts.append(f"[dim]{total - running - failed} stopped[/dim]")
|
||||
|
||||
console.print(" | ".join(summary_parts))
|
||||
console.print()
|
||||
@@ -0,0 +1,183 @@
|
||||
"""
|
||||
pyserve top - Live monitoring dashboard
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("top")
|
||||
@click.argument("services", nargs=-1)
|
||||
@click.option(
|
||||
"--refresh",
|
||||
"refresh_interval",
|
||||
default=2,
|
||||
type=float,
|
||||
help="Refresh interval in seconds",
|
||||
)
|
||||
@click.option(
|
||||
"--no-color",
|
||||
is_flag=True,
|
||||
help="Disable colored output",
|
||||
)
|
||||
@click.pass_obj
|
||||
def top_cmd(ctx: Any, services: tuple[str, ...], refresh_interval: float, no_color: bool) -> None:
|
||||
"""
|
||||
Live monitoring dashboard for services.
|
||||
|
||||
Shows real-time CPU, memory usage, and request metrics.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve top # Monitor all services
|
||||
pyserve top api admin # Monitor specific services
|
||||
pyserve top --refresh 5 # Slower refresh rate
|
||||
"""
|
||||
from ..output import console, print_info
|
||||
from ..state import StateManager
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
|
||||
if not state_manager.is_daemon_running():
|
||||
print_info("No services running. Start with 'pyserve up -d'")
|
||||
return
|
||||
|
||||
try:
|
||||
asyncio.run(
|
||||
_run_dashboard(
|
||||
state_manager,
|
||||
list(services) if services else None,
|
||||
refresh_interval,
|
||||
no_color,
|
||||
)
|
||||
)
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n")
|
||||
|
||||
|
||||
async def _run_dashboard(
|
||||
state_manager: Any,
|
||||
filter_services: Optional[list[str]],
|
||||
refresh_interval: float,
|
||||
no_color: bool,
|
||||
) -> None:
|
||||
from rich.layout import Layout
|
||||
from rich.live import Live
|
||||
from rich.panel import Panel
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
from ..output import console, format_bytes, format_uptime
|
||||
|
||||
try:
|
||||
import psutil
|
||||
except ImportError:
|
||||
console.print("[yellow]psutil not installed. Install with: pip install psutil[/yellow]")
|
||||
return
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
def make_dashboard() -> Any:
|
||||
all_services = state_manager.get_all_services()
|
||||
|
||||
if filter_services:
|
||||
all_services = {k: v for k, v in all_services.items() if k in filter_services}
|
||||
|
||||
table = Table(
|
||||
title=None,
|
||||
show_header=True,
|
||||
header_style="bold",
|
||||
border_style="dim",
|
||||
expand=True,
|
||||
)
|
||||
table.add_column("SERVICE", style="cyan", no_wrap=True)
|
||||
table.add_column("STATUS", no_wrap=True)
|
||||
table.add_column("CPU%", justify="right")
|
||||
table.add_column("MEM", justify="right")
|
||||
table.add_column("PID", style="dim")
|
||||
table.add_column("UPTIME", style="dim")
|
||||
table.add_column("HEALTH", no_wrap=True)
|
||||
|
||||
total_cpu = 0.0
|
||||
total_mem = 0
|
||||
running_count = 0
|
||||
total_count = len(all_services)
|
||||
|
||||
for name, service in sorted(all_services.items()):
|
||||
status_style = {
|
||||
"running": "[green]● RUN[/green]",
|
||||
"stopped": "[dim]○ STOP[/dim]",
|
||||
"failed": "[red]✗ FAIL[/red]",
|
||||
"starting": "[yellow]◐ START[/yellow]",
|
||||
"stopping": "[yellow]◑ STOP[/yellow]",
|
||||
}.get(service.state, service.state)
|
||||
|
||||
cpu_str = "-"
|
||||
mem_str = "-"
|
||||
|
||||
if service.pid and service.state == "running":
|
||||
try:
|
||||
proc = psutil.Process(service.pid)
|
||||
cpu = proc.cpu_percent(interval=0.1)
|
||||
mem = proc.memory_info().rss
|
||||
cpu_str = f"{cpu:.1f}%"
|
||||
mem_str = format_bytes(mem)
|
||||
total_cpu += cpu
|
||||
total_mem += mem
|
||||
running_count += 1
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||
pass
|
||||
|
||||
health_style = {
|
||||
"healthy": "[green]✓[/green]",
|
||||
"unhealthy": "[red]✗[/red]",
|
||||
"degraded": "[yellow]⚠[/yellow]",
|
||||
"unknown": "[dim]?[/dim]",
|
||||
}.get(service.health.status, "[dim]-[/dim]")
|
||||
|
||||
uptime = format_uptime(service.uptime) if service.state == "running" else "-"
|
||||
pid = str(service.pid) if service.pid else "-"
|
||||
|
||||
table.add_row(
|
||||
name,
|
||||
status_style,
|
||||
cpu_str,
|
||||
mem_str,
|
||||
pid,
|
||||
uptime,
|
||||
health_style,
|
||||
)
|
||||
|
||||
elapsed = format_uptime(time.time() - start_time)
|
||||
summary = Text()
|
||||
summary.append(f"Running: {running_count}/{total_count}", style="bold")
|
||||
summary.append(" | ")
|
||||
summary.append(f"CPU: {total_cpu:.1f}%", style="cyan")
|
||||
summary.append(" | ")
|
||||
summary.append(f"MEM: {format_bytes(total_mem)}", style="cyan")
|
||||
summary.append(" | ")
|
||||
summary.append(f"Session: {elapsed}", style="dim")
|
||||
|
||||
layout = Layout()
|
||||
layout.split_column(
|
||||
Layout(
|
||||
Panel(
|
||||
Text("PyServe Dashboard", style="bold cyan", justify="center"),
|
||||
border_style="cyan",
|
||||
),
|
||||
size=3,
|
||||
),
|
||||
Layout(table),
|
||||
Layout(Panel(summary, border_style="dim"), size=3),
|
||||
)
|
||||
|
||||
return layout
|
||||
|
||||
with Live(make_dashboard(), refresh_per_second=1 / refresh_interval, console=console) as live:
|
||||
while True:
|
||||
await asyncio.sleep(refresh_interval)
|
||||
live.update(make_dashboard())
|
||||
@@ -0,0 +1,175 @@
|
||||
"""
|
||||
pyserve up - Start all services
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import signal
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import click
|
||||
|
||||
|
||||
@click.command("up")
|
||||
@click.argument("services", nargs=-1)
|
||||
@click.option(
|
||||
"-d",
|
||||
"--detach",
|
||||
is_flag=True,
|
||||
help="Run in background (detached mode)",
|
||||
)
|
||||
@click.option(
|
||||
"--build",
|
||||
is_flag=True,
|
||||
help="Build/reload applications before starting",
|
||||
)
|
||||
@click.option(
|
||||
"--force-recreate",
|
||||
is_flag=True,
|
||||
help="Recreate services even if configuration hasn't changed",
|
||||
)
|
||||
@click.option(
|
||||
"--scale",
|
||||
"scales",
|
||||
multiple=True,
|
||||
help="Scale SERVICE to NUM workers (e.g., --scale api=4)",
|
||||
)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
"timeout",
|
||||
default=60,
|
||||
type=int,
|
||||
help="Timeout in seconds for service startup",
|
||||
)
|
||||
@click.option(
|
||||
"--wait",
|
||||
is_flag=True,
|
||||
help="Wait for services to be healthy before returning",
|
||||
)
|
||||
@click.option(
|
||||
"--remove-orphans",
|
||||
is_flag=True,
|
||||
help="Remove services not defined in configuration",
|
||||
)
|
||||
@click.pass_obj
|
||||
def up_cmd(
|
||||
ctx: Any,
|
||||
services: tuple[str, ...],
|
||||
detach: bool,
|
||||
build: bool,
|
||||
force_recreate: bool,
|
||||
scales: tuple[str, ...],
|
||||
timeout: int,
|
||||
wait: bool,
|
||||
remove_orphans: bool,
|
||||
) -> None:
|
||||
"""
|
||||
Start services defined in configuration.
|
||||
|
||||
If no services are specified, all services will be started.
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyserve up # Start all services
|
||||
pyserve up -d # Start in background
|
||||
pyserve up api admin # Start specific services
|
||||
pyserve up --scale api=4 # Scale api to 4 workers
|
||||
pyserve up --wait # Wait for healthy status
|
||||
"""
|
||||
from .._runner import ServiceRunner
|
||||
from ..output import console, print_error, print_info, print_success, print_warning
|
||||
from ..state import StateManager
|
||||
|
||||
config_path = Path(ctx.config_file)
|
||||
|
||||
if not config_path.exists():
|
||||
print_error(f"Configuration file not found: {config_path}")
|
||||
print_info("Run 'pyserve init' to create a configuration file")
|
||||
raise click.Abort()
|
||||
|
||||
scale_map = {}
|
||||
for scale in scales:
|
||||
try:
|
||||
service, num = scale.split("=")
|
||||
scale_map[service] = int(num)
|
||||
except ValueError:
|
||||
print_error(f"Invalid scale format: {scale}. Use SERVICE=NUM")
|
||||
raise click.Abort()
|
||||
|
||||
try:
|
||||
from ...config import Config
|
||||
|
||||
config = Config.from_yaml(str(config_path))
|
||||
except Exception as e:
|
||||
print_error(f"Failed to load configuration: {e}")
|
||||
raise click.Abort()
|
||||
|
||||
state_manager = StateManager(Path(".pyserve"), ctx.project)
|
||||
|
||||
if state_manager.is_daemon_running():
|
||||
daemon_pid = state_manager.get_daemon_pid()
|
||||
print_warning(f"PyServe daemon is already running (PID: {daemon_pid})")
|
||||
if not click.confirm("Do you want to restart it?"):
|
||||
raise click.Abort()
|
||||
try:
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
# FIXME: Please fix the cast usage here
|
||||
os.kill(cast(int, daemon_pid), signal.SIGTERM)
|
||||
time.sleep(2)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
state_manager.clear_daemon_pid()
|
||||
|
||||
runner = ServiceRunner(config, state_manager)
|
||||
|
||||
service_list = list(services) if services else None
|
||||
|
||||
if detach:
|
||||
console.print("[bold]Starting PyServe in background...[/bold]")
|
||||
|
||||
try:
|
||||
pid = runner.start_daemon(
|
||||
service_list,
|
||||
scale_map=scale_map,
|
||||
force_recreate=force_recreate,
|
||||
)
|
||||
state_manager.set_daemon_pid(pid)
|
||||
print_success(f"PyServe started in background (PID: {pid})")
|
||||
print_info("Use 'pyserve ps' to see service status")
|
||||
print_info("Use 'pyserve logs -f' to follow logs")
|
||||
print_info("Use 'pyserve down' to stop")
|
||||
except Exception as e:
|
||||
print_error(f"Failed to start daemon: {e}")
|
||||
raise click.Abort()
|
||||
else:
|
||||
console.print("[bold]Starting PyServe...[/bold]")
|
||||
|
||||
def signal_handler(signum: int, frame: Any) -> None:
|
||||
console.print("\n[yellow]Received shutdown signal...[/yellow]")
|
||||
runner.stop()
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
try:
|
||||
asyncio.run(
|
||||
runner.start(
|
||||
service_list,
|
||||
scale_map=scale_map,
|
||||
force_recreate=force_recreate,
|
||||
wait_healthy=wait,
|
||||
timeout=timeout,
|
||||
)
|
||||
)
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n[yellow]Shutting down...[/yellow]")
|
||||
except Exception as e:
|
||||
print_error(f"Failed to start services: {e}")
|
||||
if ctx.debug:
|
||||
raise
|
||||
raise click.Abort()
|
||||
@@ -0,0 +1,168 @@
|
||||
"""
|
||||
PyServeCTL - Main entry point
|
||||
|
||||
Usage:
|
||||
pyservectl [OPTIONS] COMMAND [ARGS]...
|
||||
"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import click
|
||||
|
||||
from .. import __version__
|
||||
from .commands import (
|
||||
config_cmd,
|
||||
down_cmd,
|
||||
health_cmd,
|
||||
init_cmd,
|
||||
logs_cmd,
|
||||
ps_cmd,
|
||||
restart_cmd,
|
||||
scale_cmd,
|
||||
start_cmd,
|
||||
stop_cmd,
|
||||
top_cmd,
|
||||
up_cmd,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config import Config
|
||||
from .state import StateManager
|
||||
|
||||
DEFAULT_CONFIG = "config.yaml"
|
||||
DEFAULT_STATE_DIR = ".pyserve"
|
||||
|
||||
|
||||
class Context:
|
||||
def __init__(self) -> None:
|
||||
self.config_file: str = DEFAULT_CONFIG
|
||||
self.state_dir: Path = Path(DEFAULT_STATE_DIR)
|
||||
self.verbose: bool = False
|
||||
self.debug: bool = False
|
||||
self.project: Optional[str] = None
|
||||
self._config: Optional["Config"] = None
|
||||
self._state: Optional["StateManager"] = None
|
||||
|
||||
@property
|
||||
def config(self) -> "Config":
|
||||
if self._config is None:
|
||||
from ..config import Config
|
||||
|
||||
if Path(self.config_file).exists():
|
||||
self._config = Config.from_yaml(self.config_file)
|
||||
else:
|
||||
self._config = Config()
|
||||
return self._config
|
||||
|
||||
@property
|
||||
def state(self) -> "StateManager":
|
||||
if self._state is None:
|
||||
from .state import StateManager
|
||||
|
||||
self._state = StateManager(self.state_dir, self.project)
|
||||
return self._state
|
||||
|
||||
|
||||
pass_context = click.make_pass_decorator(Context, ensure=True)
|
||||
|
||||
|
||||
@click.group(invoke_without_command=True)
|
||||
@click.option(
|
||||
"-c",
|
||||
"--config",
|
||||
"config_file",
|
||||
default=DEFAULT_CONFIG,
|
||||
envvar="PYSERVE_CONFIG",
|
||||
help=f"Path to configuration file (default: {DEFAULT_CONFIG})",
|
||||
type=click.Path(),
|
||||
)
|
||||
@click.option(
|
||||
"-p",
|
||||
"--project",
|
||||
"project",
|
||||
default=None,
|
||||
envvar="PYSERVE_PROJECT",
|
||||
help="Project name for isolation",
|
||||
)
|
||||
@click.option(
|
||||
"-v",
|
||||
"--verbose",
|
||||
is_flag=True,
|
||||
help="Enable verbose output",
|
||||
)
|
||||
@click.option(
|
||||
"--debug",
|
||||
is_flag=True,
|
||||
help="Enable debug mode",
|
||||
)
|
||||
@click.version_option(version=__version__, prog_name="pyservectl")
|
||||
@click.pass_context
|
||||
def cli(ctx: click.Context, config_file: str, project: Optional[str], verbose: bool, debug: bool) -> None:
|
||||
"""
|
||||
PyServeCTL - Service management CLI for PyServe.
|
||||
|
||||
Docker-compose-like tool for managing PyServe services.
|
||||
|
||||
\b
|
||||
Quick Start:
|
||||
pyservectl init # Initialize a new project
|
||||
pyservectl up # Start all services
|
||||
pyservectl ps # Show service status
|
||||
pyservectl logs -f # Follow logs
|
||||
pyservectl down # Stop all services
|
||||
|
||||
\b
|
||||
Examples:
|
||||
pyservectl up -d # Start in background
|
||||
pyservectl up -c prod.yaml # Use custom config
|
||||
pyservectl logs api -f --tail 100 # Follow api logs
|
||||
pyservectl restart api admin # Restart specific services
|
||||
pyservectl scale api=4 # Scale api to 4 workers
|
||||
"""
|
||||
ctx.ensure_object(Context)
|
||||
ctx.obj.config_file = config_file
|
||||
ctx.obj.verbose = verbose
|
||||
ctx.obj.debug = debug
|
||||
ctx.obj.project = project
|
||||
|
||||
if ctx.invoked_subcommand is None:
|
||||
click.echo(ctx.get_help())
|
||||
|
||||
|
||||
cli.add_command(init_cmd, name="init")
|
||||
cli.add_command(config_cmd, name="config")
|
||||
cli.add_command(up_cmd, name="up")
|
||||
cli.add_command(down_cmd, name="down")
|
||||
cli.add_command(start_cmd, name="start")
|
||||
cli.add_command(stop_cmd, name="stop")
|
||||
cli.add_command(restart_cmd, name="restart")
|
||||
cli.add_command(ps_cmd, name="ps")
|
||||
cli.add_command(logs_cmd, name="logs")
|
||||
cli.add_command(top_cmd, name="top")
|
||||
cli.add_command(health_cmd, name="health")
|
||||
cli.add_command(scale_cmd, name="scale")
|
||||
|
||||
# Alias 'status' -> 'ps'
|
||||
cli.add_command(ps_cmd, name="status")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
try:
|
||||
cli(standalone_mode=False)
|
||||
except click.ClickException as e:
|
||||
e.show()
|
||||
sys.exit(e.exit_code)
|
||||
except KeyboardInterrupt:
|
||||
click.echo("\nInterrupted by user")
|
||||
sys.exit(130)
|
||||
except Exception as e:
|
||||
if "--debug" in sys.argv:
|
||||
raise
|
||||
click.secho(f"Error: {e}", fg="red", err=True)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,110 @@
|
||||
"""
|
||||
PyServe CLI Output utilities
|
||||
|
||||
Rich-based formatters and helpers for CLI output.
|
||||
"""
|
||||
|
||||
from rich.console import Console
|
||||
from rich.table import Table
|
||||
from rich.theme import Theme
|
||||
|
||||
pyserve_theme = Theme(
|
||||
{
|
||||
"info": "cyan",
|
||||
"warning": "yellow",
|
||||
"error": "red bold",
|
||||
"success": "green",
|
||||
"service.running": "green",
|
||||
"service.stopped": "dim",
|
||||
"service.failed": "red",
|
||||
"service.starting": "yellow",
|
||||
}
|
||||
)
|
||||
|
||||
console = Console(theme=pyserve_theme)
|
||||
|
||||
|
||||
def print_error(message: str) -> None:
|
||||
console.print(f"[error] {message}[/error]")
|
||||
|
||||
|
||||
def print_warning(message: str) -> None:
|
||||
console.print(f"[warning] {message}[/warning]")
|
||||
|
||||
|
||||
def print_success(message: str) -> None:
|
||||
console.print(f"[success] {message}[/success]")
|
||||
|
||||
|
||||
def print_info(message: str) -> None:
|
||||
console.print(f"[info] {message}[/info]")
|
||||
|
||||
|
||||
def create_services_table() -> Table:
|
||||
table = Table(
|
||||
title=None,
|
||||
show_header=True,
|
||||
header_style="bold",
|
||||
border_style="dim",
|
||||
)
|
||||
table.add_column("NAME", style="cyan", no_wrap=True)
|
||||
table.add_column("STATUS", no_wrap=True)
|
||||
table.add_column("PORTS", style="dim")
|
||||
table.add_column("UPTIME", style="dim")
|
||||
table.add_column("HEALTH", no_wrap=True)
|
||||
table.add_column("PID", style="dim")
|
||||
table.add_column("WORKERS", style="dim")
|
||||
return table
|
||||
|
||||
|
||||
def format_status(status: str) -> str:
|
||||
status_styles = {
|
||||
"running": "[service.running]● running[/service.running]",
|
||||
"stopped": "[service.stopped]○ stopped[/service.stopped]",
|
||||
"failed": "[service.failed]✗ failed[/service.failed]",
|
||||
"starting": "[service.starting]◐ starting[/service.starting]",
|
||||
"stopping": "[service.starting]◑ stopping[/service.starting]",
|
||||
"restarting": "[service.starting]↻ restarting[/service.starting]",
|
||||
"pending": "[service.stopped]○ pending[/service.stopped]",
|
||||
}
|
||||
return status_styles.get(status.lower(), status)
|
||||
|
||||
|
||||
def format_health(health: str) -> str:
|
||||
health_styles = {
|
||||
"healthy": "[green] healthy[/green]",
|
||||
"unhealthy": "[red] unhealthy[/red]",
|
||||
"degraded": "[yellow] degraded[/yellow]",
|
||||
"unknown": "[dim] unknown[/dim]",
|
||||
"-": "[dim]-[/dim]",
|
||||
}
|
||||
return health_styles.get(health.lower(), health)
|
||||
|
||||
|
||||
def format_uptime(seconds: float) -> str:
|
||||
if seconds <= 0:
|
||||
return "-"
|
||||
|
||||
if seconds < 60:
|
||||
return f"{int(seconds)}s"
|
||||
elif seconds < 3600:
|
||||
minutes = int(seconds / 60)
|
||||
secs = int(seconds % 60)
|
||||
return f"{minutes}m {secs}s"
|
||||
elif seconds < 86400:
|
||||
hours = int(seconds / 3600)
|
||||
minutes = int((seconds % 3600) / 60)
|
||||
return f"{hours}h {minutes}m"
|
||||
else:
|
||||
days = int(seconds / 86400)
|
||||
hours = int((seconds % 86400) / 3600)
|
||||
return f"{days}d {hours}h"
|
||||
|
||||
|
||||
def format_bytes(num_bytes: int) -> str:
|
||||
value = float(num_bytes)
|
||||
for unit in ["B", "KB", "MB", "GB", "TB"]:
|
||||
if abs(value) < 1024.0:
|
||||
return f"{value:.1f}{unit}"
|
||||
value /= 1024.0
|
||||
return f"{value:.1f}PB"
|
||||
@@ -0,0 +1,232 @@
|
||||
"""
|
||||
PyServe CLI State Management
|
||||
|
||||
Manages the state of running services.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServiceHealth:
|
||||
status: str = "unknown" # healthy, unhealthy, degraded, unknown
|
||||
last_check: Optional[float] = None
|
||||
failures: int = 0
|
||||
response_time_ms: Optional[float] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServiceState:
|
||||
name: str
|
||||
state: str = "stopped" # pending, starting, running, stopping, stopped, failed, restarting
|
||||
pid: Optional[int] = None
|
||||
port: int = 0
|
||||
workers: int = 0
|
||||
started_at: Optional[float] = None
|
||||
restart_count: int = 0
|
||||
health: ServiceHealth = field(default_factory=ServiceHealth)
|
||||
config_hash: str = ""
|
||||
|
||||
@property
|
||||
def uptime(self) -> float:
|
||||
if self.started_at is None:
|
||||
return 0.0
|
||||
return time.time() - self.started_at
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"state": self.state,
|
||||
"pid": self.pid,
|
||||
"port": self.port,
|
||||
"workers": self.workers,
|
||||
"started_at": self.started_at,
|
||||
"restart_count": self.restart_count,
|
||||
"health": asdict(self.health),
|
||||
"config_hash": self.config_hash,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> "ServiceState":
|
||||
health_data = data.pop("health", {})
|
||||
health = ServiceHealth(**health_data) if health_data else ServiceHealth()
|
||||
return cls(**data, health=health)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProjectState:
|
||||
version: str = "1.0"
|
||||
project: str = ""
|
||||
config_file: str = ""
|
||||
config_hash: str = ""
|
||||
started_at: Optional[float] = None
|
||||
daemon_pid: Optional[int] = None
|
||||
services: Dict[str, ServiceState] = field(default_factory=dict)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"version": self.version,
|
||||
"project": self.project,
|
||||
"config_file": self.config_file,
|
||||
"config_hash": self.config_hash,
|
||||
"started_at": self.started_at,
|
||||
"daemon_pid": self.daemon_pid,
|
||||
"services": {name: svc.to_dict() for name, svc in self.services.items()},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> "ProjectState":
|
||||
services_data = data.pop("services", {})
|
||||
services = {name: ServiceState.from_dict(svc) for name, svc in services_data.items()}
|
||||
return cls(**data, services=services)
|
||||
|
||||
|
||||
class StateManager:
|
||||
STATE_FILE = "state.json"
|
||||
PID_FILE = "pyserve.pid"
|
||||
SOCKET_FILE = "pyserve.sock"
|
||||
LOGS_DIR = "logs"
|
||||
|
||||
def __init__(self, state_dir: Path, project: Optional[str] = None):
|
||||
self.state_dir = Path(state_dir)
|
||||
self.project = project or self._detect_project()
|
||||
self._state: Optional[ProjectState] = None
|
||||
|
||||
def _detect_project(self) -> str:
|
||||
return Path.cwd().name
|
||||
|
||||
@property
|
||||
def state_file(self) -> Path:
|
||||
return self.state_dir / self.STATE_FILE
|
||||
|
||||
@property
|
||||
def pid_file(self) -> Path:
|
||||
return self.state_dir / self.PID_FILE
|
||||
|
||||
@property
|
||||
def socket_file(self) -> Path:
|
||||
return self.state_dir / self.SOCKET_FILE
|
||||
|
||||
@property
|
||||
def logs_dir(self) -> Path:
|
||||
return self.state_dir / self.LOGS_DIR
|
||||
|
||||
def ensure_dirs(self) -> None:
|
||||
self.state_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.logs_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def load(self) -> ProjectState:
|
||||
if self._state is not None:
|
||||
return self._state
|
||||
|
||||
if self.state_file.exists():
|
||||
try:
|
||||
with open(self.state_file) as f:
|
||||
data = json.load(f)
|
||||
self._state = ProjectState.from_dict(data)
|
||||
except (json.JSONDecodeError, KeyError):
|
||||
self._state = ProjectState(project=self.project)
|
||||
else:
|
||||
self._state = ProjectState(project=self.project)
|
||||
|
||||
return self._state
|
||||
|
||||
def save(self) -> None:
|
||||
if self._state is None:
|
||||
return
|
||||
|
||||
self.ensure_dirs()
|
||||
|
||||
with open(self.state_file, "w") as f:
|
||||
json.dump(self._state.to_dict(), f, indent=2)
|
||||
|
||||
def get_state(self) -> ProjectState:
|
||||
return self.load()
|
||||
|
||||
def update_service(self, name: str, **kwargs: Any) -> ServiceState:
|
||||
state = self.load()
|
||||
|
||||
if name not in state.services:
|
||||
state.services[name] = ServiceState(name=name)
|
||||
|
||||
service = state.services[name]
|
||||
for key, value in kwargs.items():
|
||||
if hasattr(service, key):
|
||||
setattr(service, key, value)
|
||||
|
||||
self.save()
|
||||
return service
|
||||
|
||||
def remove_service(self, name: str) -> None:
|
||||
state = self.load()
|
||||
if name in state.services:
|
||||
del state.services[name]
|
||||
self.save()
|
||||
|
||||
def get_service(self, name: str) -> Optional[ServiceState]:
|
||||
state = self.load()
|
||||
return state.services.get(name)
|
||||
|
||||
def get_all_services(self) -> Dict[str, ServiceState]:
|
||||
state = self.load()
|
||||
return state.services.copy()
|
||||
|
||||
def clear(self) -> None:
|
||||
self._state = ProjectState(project=self.project)
|
||||
self.save()
|
||||
|
||||
def is_daemon_running(self) -> bool:
|
||||
if not self.pid_file.exists():
|
||||
return False
|
||||
|
||||
try:
|
||||
pid = int(self.pid_file.read_text().strip())
|
||||
# Check if process exists
|
||||
os.kill(pid, 0)
|
||||
return True
|
||||
except (ValueError, ProcessLookupError, PermissionError):
|
||||
return False
|
||||
|
||||
def get_daemon_pid(self) -> Optional[int]:
|
||||
if not self.is_daemon_running():
|
||||
return None
|
||||
|
||||
try:
|
||||
return int(self.pid_file.read_text().strip())
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def set_daemon_pid(self, pid: int) -> None:
|
||||
self.ensure_dirs()
|
||||
self.pid_file.write_text(str(pid))
|
||||
|
||||
state = self.load()
|
||||
state.daemon_pid = pid
|
||||
self.save()
|
||||
|
||||
def clear_daemon_pid(self) -> None:
|
||||
if self.pid_file.exists():
|
||||
self.pid_file.unlink()
|
||||
|
||||
state = self.load()
|
||||
state.daemon_pid = None
|
||||
self.save()
|
||||
|
||||
def get_service_log_file(self, service_name: str) -> Path:
|
||||
self.ensure_dirs()
|
||||
return self.logs_dir / f"{service_name}.log"
|
||||
|
||||
def compute_config_hash(self, config_file: str) -> str:
|
||||
import hashlib
|
||||
|
||||
path = Path(config_file)
|
||||
if not path.exists():
|
||||
return ""
|
||||
|
||||
content = path.read_bytes()
|
||||
return hashlib.sha256(content).hexdigest()[:16]
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
|
||||
@@ -216,6 +217,15 @@ class ExtensionManager:
|
||||
"monitoring": MonitoringExtension,
|
||||
"asgi": ASGIExtension,
|
||||
}
|
||||
self._register_process_orchestration()
|
||||
|
||||
def _register_process_orchestration(self) -> None:
|
||||
try:
|
||||
from .process_extension import ProcessOrchestrationExtension
|
||||
|
||||
self.extension_registry["process_orchestration"] = ProcessOrchestrationExtension
|
||||
except ImportError:
|
||||
pass # Optional dependency
|
||||
|
||||
def register_extension_type(self, name: str, extension_class: Type[Extension]) -> None:
|
||||
self.extension_registry[name] = extension_class
|
||||
@@ -234,6 +244,32 @@ class ExtensionManager:
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading extension {extension_type}: {e}")
|
||||
|
||||
async def load_extension_async(self, extension_type: str, config: Dict[str, Any]) -> None:
|
||||
"""Load extension with async setup support (for ProcessOrchestration)."""
|
||||
if extension_type not in self.extension_registry:
|
||||
logger.error(f"Unknown extension type: {extension_type}")
|
||||
return
|
||||
|
||||
try:
|
||||
extension_class = self.extension_registry[extension_type]
|
||||
extension = extension_class(config)
|
||||
|
||||
setup_method = getattr(extension, "setup", None)
|
||||
if setup_method is not None and asyncio.iscoroutinefunction(setup_method):
|
||||
await setup_method(config)
|
||||
else:
|
||||
extension.initialize()
|
||||
|
||||
start_method = getattr(extension, "start", None)
|
||||
if start_method is not None and asyncio.iscoroutinefunction(start_method):
|
||||
await start_method()
|
||||
|
||||
# Insert at the beginning so process_orchestration is checked first
|
||||
self.extensions.insert(0, extension)
|
||||
logger.info(f"Loaded extension (async): {extension_type}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading extension {extension_type}: {e}")
|
||||
|
||||
async def process_request(self, request: Request) -> Optional[Response]:
|
||||
for extension in self.extensions:
|
||||
if not extension.enabled:
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
"""Process Orchestration Extension
|
||||
|
||||
Extension that manages ASGI/WSGI applications as isolated processes
|
||||
and routes requests to them via reverse proxy.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from .extensions import Extension
|
||||
from .logging_utils import get_logger
|
||||
from .process_manager import ProcessConfig, ProcessManager
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ProcessOrchestrationExtension(Extension):
|
||||
"""
|
||||
Extension that orchestrates ASGI/WSGI applications as separate processes.
|
||||
|
||||
Unlike ASGIExtension which runs apps in-process, this extension:
|
||||
- Runs each app in its own isolated process
|
||||
- Provides health monitoring and auto-restart
|
||||
- Routes requests via HTTP reverse proxy
|
||||
- Supports multiple workers per app
|
||||
|
||||
Configuration example:
|
||||
```yaml
|
||||
extensions:
|
||||
- type: process_orchestration
|
||||
config:
|
||||
port_range: [9000, 9999]
|
||||
health_check_enabled: true
|
||||
apps:
|
||||
- name: api
|
||||
path: /api
|
||||
app_path: myapp.api:app
|
||||
workers: 4
|
||||
health_check_path: /health
|
||||
|
||||
- name: admin
|
||||
path: /admin
|
||||
app_path: myapp.admin:create_app
|
||||
factory: true
|
||||
workers: 2
|
||||
```
|
||||
"""
|
||||
|
||||
name = "process_orchestration"
|
||||
|
||||
def __init__(self, config: Dict[str, Any]) -> None:
|
||||
super().__init__(config)
|
||||
self._manager: Optional[ProcessManager] = None
|
||||
self._mounts: Dict[str, MountConfig] = {} # path -> config
|
||||
self._http_client: Optional[httpx.AsyncClient] = None
|
||||
self._started = False
|
||||
self._proxy_timeout: float = config.get("proxy_timeout", 60.0)
|
||||
self._pending_config = config # Store for async setup
|
||||
|
||||
logging_config = config.get("logging", {})
|
||||
self._log_proxy_requests: bool = logging_config.get("proxy_logs", True)
|
||||
self._log_health_checks: bool = logging_config.get("health_check_logs", False)
|
||||
|
||||
httpx_level = logging_config.get("httpx_level", "warning").upper()
|
||||
logging.getLogger("httpx").setLevel(getattr(logging, httpx_level, logging.WARNING))
|
||||
logging.getLogger("httpcore").setLevel(getattr(logging, httpx_level, logging.WARNING))
|
||||
|
||||
async def setup(self, config: Optional[Dict[str, Any]] = None) -> None:
|
||||
if config is None:
|
||||
config = self._pending_config
|
||||
|
||||
port_range = tuple(config.get("port_range", [9000, 9999]))
|
||||
health_check_enabled = config.get("health_check_enabled", True)
|
||||
self._proxy_timeout = config.get("proxy_timeout", 60.0)
|
||||
|
||||
self._manager = ProcessManager(
|
||||
port_range=port_range,
|
||||
health_check_enabled=health_check_enabled,
|
||||
)
|
||||
|
||||
self._http_client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(self._proxy_timeout),
|
||||
follow_redirects=False,
|
||||
limits=httpx.Limits(
|
||||
max_keepalive_connections=100,
|
||||
max_connections=200,
|
||||
),
|
||||
)
|
||||
|
||||
apps_config = config.get("apps", [])
|
||||
for app_config in apps_config:
|
||||
await self._register_app(app_config)
|
||||
|
||||
logger.info(
|
||||
"Process orchestration extension initialized",
|
||||
app_count=len(self._mounts),
|
||||
)
|
||||
|
||||
async def _register_app(self, app_config: Dict[str, Any]) -> None:
|
||||
if not self._manager:
|
||||
return
|
||||
|
||||
name = app_config.get("name")
|
||||
path = app_config.get("path", "").rstrip("/")
|
||||
app_path = app_config.get("app_path")
|
||||
|
||||
if not name or not app_path:
|
||||
logger.error("App config missing 'name' or 'app_path'")
|
||||
return
|
||||
|
||||
process_config = ProcessConfig(
|
||||
name=name,
|
||||
app_path=app_path,
|
||||
app_type=app_config.get("app_type", "asgi"),
|
||||
workers=app_config.get("workers", 1),
|
||||
module_path=app_config.get("module_path"),
|
||||
factory=app_config.get("factory", False),
|
||||
factory_args=app_config.get("factory_args"),
|
||||
env=app_config.get("env", {}),
|
||||
health_check_enabled=app_config.get("health_check_enabled", True),
|
||||
health_check_path=app_config.get("health_check_path", "/health"),
|
||||
health_check_interval=app_config.get("health_check_interval", 10.0),
|
||||
health_check_timeout=app_config.get("health_check_timeout", 5.0),
|
||||
health_check_retries=app_config.get("health_check_retries", 3),
|
||||
max_memory_mb=app_config.get("max_memory_mb"),
|
||||
max_restart_count=app_config.get("max_restart_count", 5),
|
||||
restart_delay=app_config.get("restart_delay", 1.0),
|
||||
shutdown_timeout=app_config.get("shutdown_timeout", 30.0),
|
||||
)
|
||||
|
||||
await self._manager.register(process_config)
|
||||
|
||||
self._mounts[path] = MountConfig(
|
||||
path=path,
|
||||
process_name=name,
|
||||
strip_path=app_config.get("strip_path", True),
|
||||
)
|
||||
|
||||
logger.info(f"Registered app '{name}' at path '{path}'")
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._started or not self._manager:
|
||||
return
|
||||
|
||||
await self._manager.start()
|
||||
results = await self._manager.start_all()
|
||||
|
||||
self._started = True
|
||||
|
||||
success = sum(1 for v in results.values() if v)
|
||||
failed = len(results) - success
|
||||
|
||||
logger.info(
|
||||
"Process orchestration started",
|
||||
success=success,
|
||||
failed=failed,
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._started:
|
||||
return
|
||||
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
if self._manager:
|
||||
await self._manager.stop()
|
||||
|
||||
self._started = False
|
||||
logger.info("Process orchestration stopped")
|
||||
|
||||
def cleanup(self) -> None:
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
loop.create_task(self.stop())
|
||||
except RuntimeError:
|
||||
asyncio.run(self.stop())
|
||||
|
||||
async def process_request(self, request: Request) -> Optional[Response]:
|
||||
if not self._started or not self._manager:
|
||||
logger.debug(
|
||||
"Process orchestration not ready",
|
||||
started=self._started,
|
||||
has_manager=self._manager is not None,
|
||||
)
|
||||
return None
|
||||
|
||||
mount = self._get_mount(request.url.path)
|
||||
if not mount:
|
||||
logger.debug(
|
||||
"No mount found for path",
|
||||
path=request.url.path,
|
||||
available_mounts=list(self._mounts.keys()),
|
||||
)
|
||||
return None
|
||||
|
||||
upstream_url = self._manager.get_upstream_url(mount.process_name)
|
||||
if not upstream_url:
|
||||
logger.warning(
|
||||
f"Process '{mount.process_name}' not running",
|
||||
path=request.url.path,
|
||||
)
|
||||
return Response("Service Unavailable", status_code=503)
|
||||
|
||||
request_id = request.headers.get("X-Request-ID", str(uuid.uuid4())[:8])
|
||||
|
||||
start_time = time.perf_counter()
|
||||
response = await self._proxy_request(request, upstream_url, mount, request_id)
|
||||
latency_ms = (time.perf_counter() - start_time) * 1000
|
||||
|
||||
if self._log_proxy_requests:
|
||||
logger.info(
|
||||
"Proxy request completed",
|
||||
request_id=request_id,
|
||||
method=request.method,
|
||||
path=request.url.path,
|
||||
process=mount.process_name,
|
||||
upstream=upstream_url,
|
||||
status=response.status_code,
|
||||
latency_ms=round(latency_ms, 2),
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
def _get_mount(self, path: str) -> Optional["MountConfig"]:
|
||||
for mount_path in sorted(self._mounts.keys(), key=len, reverse=True):
|
||||
if mount_path == "":
|
||||
return self._mounts[mount_path]
|
||||
if path == mount_path or path.startswith(f"{mount_path}/"):
|
||||
return self._mounts[mount_path]
|
||||
return None
|
||||
|
||||
async def _proxy_request(
|
||||
self,
|
||||
request: Request,
|
||||
upstream_url: str,
|
||||
mount: "MountConfig",
|
||||
request_id: str = "",
|
||||
) -> Response:
|
||||
path = request.url.path
|
||||
if mount.strip_path and mount.path:
|
||||
path = path[len(mount.path) :] or "/"
|
||||
|
||||
target_url = f"{upstream_url}{path}"
|
||||
if request.url.query:
|
||||
target_url += f"?{request.url.query}"
|
||||
|
||||
headers = dict(request.headers)
|
||||
headers.pop("host", None)
|
||||
headers["X-Forwarded-For"] = request.client.host if request.client else "unknown"
|
||||
headers["X-Forwarded-Proto"] = request.url.scheme
|
||||
headers["X-Forwarded-Host"] = request.headers.get("host", "")
|
||||
if request_id:
|
||||
headers["X-Request-ID"] = request_id
|
||||
|
||||
try:
|
||||
if not self._http_client:
|
||||
return Response("Service Unavailable", status_code=503)
|
||||
|
||||
body = await request.body()
|
||||
|
||||
response = await self._http_client.request(
|
||||
method=request.method,
|
||||
url=target_url,
|
||||
headers=headers,
|
||||
content=body,
|
||||
)
|
||||
|
||||
response_headers = dict(response.headers)
|
||||
for header in ["transfer-encoding", "connection", "keep-alive"]:
|
||||
response_headers.pop(header, None)
|
||||
|
||||
return Response(
|
||||
content=response.content,
|
||||
status_code=response.status_code,
|
||||
headers=response_headers,
|
||||
)
|
||||
|
||||
except httpx.TimeoutException:
|
||||
logger.error(f"Proxy timeout to {upstream_url}")
|
||||
return Response("Gateway Timeout", status_code=504)
|
||||
except httpx.ConnectError as e:
|
||||
logger.error(f"Proxy connection error to {upstream_url}: {e}")
|
||||
return Response("Bad Gateway", status_code=502)
|
||||
except Exception as e:
|
||||
logger.error(f"Proxy error to {upstream_url}: {e}")
|
||||
return Response("Internal Server Error", status_code=500)
|
||||
|
||||
async def process_response(
|
||||
self,
|
||||
request: Request,
|
||||
response: Response,
|
||||
) -> Response:
|
||||
return response
|
||||
|
||||
def get_metrics(self) -> Dict[str, Any]:
|
||||
metrics = {
|
||||
"process_orchestration": {
|
||||
"enabled": self._started,
|
||||
"mounts": len(self._mounts),
|
||||
}
|
||||
}
|
||||
|
||||
if self._manager:
|
||||
metrics["process_orchestration"].update(self._manager.get_metrics())
|
||||
|
||||
return metrics
|
||||
|
||||
async def get_process_status(self, name: str) -> Optional[Dict[str, Any]]:
|
||||
if not self._manager:
|
||||
return None
|
||||
info = self._manager.get_process(name)
|
||||
return info.to_dict() if info else None
|
||||
|
||||
async def get_all_status(self) -> Dict[str, Any]:
|
||||
if not self._manager:
|
||||
return {}
|
||||
return {name: info.to_dict() for name, info in self._manager.get_all_processes().items()}
|
||||
|
||||
async def restart_process(self, name: str) -> bool:
|
||||
if not self._manager:
|
||||
return False
|
||||
return await self._manager.restart_process(name)
|
||||
|
||||
async def scale_process(self, name: str, workers: int) -> bool:
|
||||
if not self._manager:
|
||||
return False
|
||||
|
||||
info = self._manager.get_process(name)
|
||||
if not info:
|
||||
return False
|
||||
|
||||
info.config.workers = workers
|
||||
return await self._manager.restart_process(name)
|
||||
|
||||
|
||||
class MountConfig:
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
process_name: str,
|
||||
strip_path: bool = True,
|
||||
):
|
||||
self.path = path
|
||||
self.process_name = process_name
|
||||
self.strip_path = strip_path
|
||||
|
||||
|
||||
async def setup_process_orchestration(config: Dict[str, Any]) -> ProcessOrchestrationExtension:
|
||||
ext = ProcessOrchestrationExtension(config)
|
||||
await ext.setup(config)
|
||||
await ext.start()
|
||||
return ext
|
||||
|
||||
|
||||
async def shutdown_process_orchestration(ext: ProcessOrchestrationExtension) -> None:
|
||||
await ext.stop()
|
||||
@@ -0,0 +1,553 @@
|
||||
"""Process Manager Module
|
||||
|
||||
Orchestrates ASGI/WSGI applications as separate processes
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from .logging_utils import get_logger
|
||||
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ProcessState(Enum):
|
||||
PENDING = "pending"
|
||||
STARTING = "starting"
|
||||
RUNNING = "running"
|
||||
STOPPING = "stopping"
|
||||
STOPPED = "stopped"
|
||||
FAILED = "failed"
|
||||
RESTARTING = "restarting"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessConfig:
|
||||
name: str
|
||||
app_path: str
|
||||
app_type: str = "asgi" # asgi, wsgi
|
||||
host: str = "127.0.0.1"
|
||||
port: int = 0 # 0 = auto-assign
|
||||
workers: int = 1
|
||||
module_path: Optional[str] = None
|
||||
factory: bool = False
|
||||
factory_args: Optional[Dict[str, Any]] = None
|
||||
env: Dict[str, str] = field(default_factory=dict)
|
||||
|
||||
health_check_enabled: bool = True
|
||||
health_check_path: str = "/health"
|
||||
health_check_interval: float = 10.0
|
||||
health_check_timeout: float = 5.0
|
||||
health_check_retries: int = 3
|
||||
|
||||
max_memory_mb: Optional[int] = None
|
||||
max_restart_count: int = 5
|
||||
restart_delay: float = 1.0 # seconds
|
||||
|
||||
shutdown_timeout: float = 30.0 # seconds
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessInfo:
|
||||
config: ProcessConfig
|
||||
state: ProcessState = ProcessState.PENDING
|
||||
pid: Optional[int] = None
|
||||
port: int = 0
|
||||
start_time: Optional[float] = None
|
||||
restart_count: int = 0
|
||||
last_health_check: Optional[float] = None
|
||||
health_check_failures: int = 0
|
||||
process: Optional[subprocess.Popen] = None
|
||||
|
||||
@property
|
||||
def uptime(self) -> float:
|
||||
if self.start_time is None:
|
||||
return 0.0
|
||||
return time.time() - self.start_time
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
return self.state == ProcessState.RUNNING and self.process is not None
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"name": self.config.name,
|
||||
"state": self.state.value,
|
||||
"pid": self.pid,
|
||||
"port": self.port,
|
||||
"uptime": round(self.uptime, 2),
|
||||
"restart_count": self.restart_count,
|
||||
"health_check_failures": self.health_check_failures,
|
||||
"workers": self.config.workers,
|
||||
}
|
||||
|
||||
|
||||
class PortAllocator:
|
||||
def __init__(self, start_port: int = 9000, end_port: int = 9999):
|
||||
self.start_port = start_port
|
||||
self.end_port = end_port
|
||||
self._allocated: set[int] = set()
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def allocate(self) -> int:
|
||||
async with self._lock:
|
||||
for port in range(self.start_port, self.end_port + 1):
|
||||
if port in self._allocated:
|
||||
continue
|
||||
if self._is_port_available(port):
|
||||
self._allocated.add(port)
|
||||
return port
|
||||
raise RuntimeError(f"No available ports in range {self.start_port}-{self.end_port}")
|
||||
|
||||
async def release(self, port: int) -> None:
|
||||
async with self._lock:
|
||||
self._allocated.discard(port)
|
||||
|
||||
def _is_port_available(self, port: int) -> bool:
|
||||
try:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
s.bind(("127.0.0.1", port))
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
class ProcessManager:
|
||||
def __init__(
|
||||
self,
|
||||
port_range: tuple[int, int] = (9000, 9999),
|
||||
health_check_enabled: bool = True,
|
||||
):
|
||||
self._processes: Dict[str, ProcessInfo] = {}
|
||||
self._port_allocator = PortAllocator(*port_range)
|
||||
self._health_check_enabled = health_check_enabled
|
||||
self._health_check_task: Optional[asyncio.Task] = None
|
||||
self._shutdown_event = asyncio.Event()
|
||||
self._started = False
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._started:
|
||||
return
|
||||
|
||||
self._started = True
|
||||
self._shutdown_event.clear()
|
||||
|
||||
if self._health_check_enabled:
|
||||
self._health_check_task = asyncio.create_task(self._health_check_loop(), name="process_manager_health_check")
|
||||
|
||||
logger.info("Process manager started")
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._started:
|
||||
return
|
||||
|
||||
logger.info("Stopping process manager...")
|
||||
self._shutdown_event.set()
|
||||
|
||||
if self._health_check_task:
|
||||
self._health_check_task.cancel()
|
||||
try:
|
||||
await self._health_check_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
await self.stop_all()
|
||||
|
||||
self._started = False
|
||||
logger.info("Process manager stopped")
|
||||
|
||||
async def register(self, config: ProcessConfig) -> ProcessInfo:
|
||||
async with self._lock:
|
||||
if config.name in self._processes:
|
||||
raise ValueError(f"Process '{config.name}' already registered")
|
||||
|
||||
info = ProcessInfo(config=config)
|
||||
self._processes[config.name] = info
|
||||
|
||||
logger.info(f"Registered process '{config.name}'", app_path=config.app_path)
|
||||
return info
|
||||
|
||||
async def unregister(self, name: str) -> None:
|
||||
async with self._lock:
|
||||
if name not in self._processes:
|
||||
return
|
||||
|
||||
info = self._processes[name]
|
||||
if info.is_running:
|
||||
await self._stop_process(info)
|
||||
|
||||
if info.port:
|
||||
await self._port_allocator.release(info.port)
|
||||
|
||||
del self._processes[name]
|
||||
logger.info(f"Unregistered process '{name}'")
|
||||
|
||||
async def start_process(self, name: str) -> bool:
|
||||
info = self._processes.get(name)
|
||||
if not info:
|
||||
logger.error(f"Process '{name}' not found")
|
||||
return False
|
||||
|
||||
if info.is_running:
|
||||
logger.warning(f"Process '{name}' is already running")
|
||||
return True
|
||||
|
||||
return await self._start_process(info)
|
||||
|
||||
async def stop_process(self, name: str) -> bool:
|
||||
info = self._processes.get(name)
|
||||
if not info:
|
||||
logger.error(f"Process '{name}' not found")
|
||||
return False
|
||||
|
||||
return await self._stop_process(info)
|
||||
|
||||
async def restart_process(self, name: str) -> bool:
|
||||
info = self._processes.get(name)
|
||||
if not info:
|
||||
logger.error(f"Process '{name}' not found")
|
||||
return False
|
||||
|
||||
info.state = ProcessState.RESTARTING
|
||||
|
||||
if info.is_running:
|
||||
await self._stop_process(info)
|
||||
|
||||
await asyncio.sleep(info.config.restart_delay)
|
||||
return await self._start_process(info)
|
||||
|
||||
async def start_all(self) -> Dict[str, bool]:
|
||||
results = {}
|
||||
for name in self._processes:
|
||||
results[name] = await self.start_process(name)
|
||||
return results
|
||||
|
||||
async def stop_all(self) -> None:
|
||||
tasks = []
|
||||
for info in self._processes.values():
|
||||
if info.is_running:
|
||||
tasks.append(self._stop_process(info))
|
||||
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
def get_process(self, name: str) -> Optional[ProcessInfo]:
|
||||
return self._processes.get(name)
|
||||
|
||||
def get_all_processes(self) -> Dict[str, ProcessInfo]:
|
||||
return self._processes.copy()
|
||||
|
||||
def get_process_by_port(self, port: int) -> Optional[ProcessInfo]:
|
||||
for info in self._processes.values():
|
||||
if info.port == port:
|
||||
return info
|
||||
return None
|
||||
|
||||
def get_upstream_url(self, name: str) -> Optional[str]:
|
||||
info = self._processes.get(name)
|
||||
if not info or not info.is_running:
|
||||
return None
|
||||
return f"http://{info.config.host}:{info.port}"
|
||||
|
||||
async def _start_process(self, info: ProcessInfo) -> bool:
|
||||
config = info.config
|
||||
|
||||
try:
|
||||
info.state = ProcessState.STARTING
|
||||
|
||||
if info.port == 0:
|
||||
info.port = await self._port_allocator.allocate()
|
||||
|
||||
cmd = self._build_command(config, info.port)
|
||||
|
||||
env = os.environ.copy()
|
||||
env.update(config.env)
|
||||
|
||||
if config.module_path:
|
||||
python_path = env.get("PYTHONPATH", "")
|
||||
module_dir = str(Path(config.module_path).resolve())
|
||||
env["PYTHONPATH"] = f"{module_dir}:{python_path}" if python_path else module_dir
|
||||
|
||||
# For WSGI apps, pass configuration via environment variables
|
||||
if config.app_type == "wsgi":
|
||||
env["PYSERVE_WSGI_APP"] = config.app_path
|
||||
env["PYSERVE_WSGI_FACTORY"] = "1" if config.factory else "0"
|
||||
|
||||
logger.info(
|
||||
f"Starting process '{config.name}'",
|
||||
command=" ".join(cmd),
|
||||
port=info.port,
|
||||
)
|
||||
|
||||
info.process = subprocess.Popen(
|
||||
cmd,
|
||||
env=env,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
preexec_fn=os.setsid if hasattr(os, "setsid") else None,
|
||||
)
|
||||
|
||||
info.pid = info.process.pid
|
||||
info.start_time = time.time()
|
||||
|
||||
if not await self._wait_for_ready(info):
|
||||
raise RuntimeError(f"Process '{config.name}' failed to start")
|
||||
|
||||
info.state = ProcessState.RUNNING
|
||||
logger.info(
|
||||
f"Process '{config.name}' started successfully",
|
||||
pid=info.pid,
|
||||
port=info.port,
|
||||
)
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to start process '{config.name}': {e}")
|
||||
info.state = ProcessState.FAILED
|
||||
if info.port:
|
||||
await self._port_allocator.release(info.port)
|
||||
info.port = 0
|
||||
return False
|
||||
|
||||
async def _stop_process(self, info: ProcessInfo) -> bool:
|
||||
if not info.process:
|
||||
info.state = ProcessState.STOPPED
|
||||
return True
|
||||
|
||||
config = info.config
|
||||
info.state = ProcessState.STOPPING
|
||||
|
||||
try:
|
||||
if hasattr(os, "killpg"):
|
||||
try:
|
||||
os.killpg(os.getpgid(info.process.pid), signal.SIGTERM)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
else:
|
||||
info.process.terminate()
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.get_event_loop().run_in_executor(None, info.process.wait), timeout=config.shutdown_timeout)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f"Process '{config.name}' did not stop gracefully, forcing kill")
|
||||
if hasattr(os, "killpg"):
|
||||
try:
|
||||
os.killpg(os.getpgid(info.process.pid), signal.SIGKILL)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
else:
|
||||
info.process.kill()
|
||||
info.process.wait()
|
||||
|
||||
if info.port:
|
||||
await self._port_allocator.release(info.port)
|
||||
|
||||
info.state = ProcessState.STOPPED
|
||||
info.process = None
|
||||
info.pid = None
|
||||
|
||||
logger.info(f"Process '{config.name}' stopped")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error stopping process '{config.name}': {e}")
|
||||
info.state = ProcessState.FAILED
|
||||
return False
|
||||
|
||||
async def _wait_for_ready(self, info: ProcessInfo, timeout: float = 30.0) -> bool:
|
||||
import httpx
|
||||
|
||||
start_time = time.time()
|
||||
url = f"http://{info.config.host}:{info.port}{info.config.health_check_path}"
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
if info.process and info.process.poll() is not None:
|
||||
stdout, stderr = info.process.communicate()
|
||||
logger.error(
|
||||
f"Process '{info.config.name}' exited during startup",
|
||||
returncode=info.process.returncode,
|
||||
stderr=stderr.decode() if stderr else "",
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=2.0) as client:
|
||||
resp = await client.get(url)
|
||||
if resp.status_code < 500:
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
return False
|
||||
|
||||
async def _health_check_loop(self) -> None:
|
||||
while not self._shutdown_event.is_set():
|
||||
try:
|
||||
for info in list(self._processes.values()):
|
||||
if not info.is_running or not info.config.health_check_enabled:
|
||||
continue
|
||||
|
||||
await self._check_process_health(info)
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._shutdown_event.wait(),
|
||||
timeout=(
|
||||
min(p.config.health_check_interval for p in self._processes.values() if p.config.health_check_enabled)
|
||||
if self._processes
|
||||
else 10.0
|
||||
),
|
||||
)
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in health check loop: {e}")
|
||||
await asyncio.sleep(5)
|
||||
|
||||
async def _check_process_health(self, info: ProcessInfo) -> bool:
|
||||
import httpx
|
||||
|
||||
config = info.config
|
||||
url = f"http://{config.host}:{info.port}{config.health_check_path}"
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=config.health_check_timeout) as client:
|
||||
resp = await client.get(url)
|
||||
if resp.status_code < 500:
|
||||
info.health_check_failures = 0
|
||||
info.last_health_check = time.time()
|
||||
return True
|
||||
else:
|
||||
raise Exception(f"Health check returned status {resp.status_code}")
|
||||
|
||||
except Exception as e:
|
||||
info.health_check_failures += 1
|
||||
logger.warning(
|
||||
f"Health check failed for '{config.name}'",
|
||||
failures=info.health_check_failures,
|
||||
error=str(e),
|
||||
)
|
||||
|
||||
if info.health_check_failures >= config.health_check_retries:
|
||||
logger.error(f"Process '{config.name}' is unhealthy, restarting...")
|
||||
await self._handle_unhealthy_process(info)
|
||||
|
||||
return False
|
||||
|
||||
async def _handle_unhealthy_process(self, info: ProcessInfo) -> None:
|
||||
config = info.config
|
||||
|
||||
if info.restart_count >= config.max_restart_count:
|
||||
logger.error(f"Process '{config.name}' exceeded max restart count, marking as failed")
|
||||
info.state = ProcessState.FAILED
|
||||
return
|
||||
|
||||
info.restart_count += 1
|
||||
info.health_check_failures = 0
|
||||
|
||||
delay = config.restart_delay * (2 ** (info.restart_count - 1))
|
||||
delay = min(delay, 60.0)
|
||||
|
||||
logger.info(
|
||||
f"Restarting process '{config.name}'",
|
||||
restart_count=info.restart_count,
|
||||
delay=delay,
|
||||
)
|
||||
|
||||
await self._stop_process(info)
|
||||
await asyncio.sleep(delay)
|
||||
await self._start_process(info)
|
||||
|
||||
def _build_command(self, config: ProcessConfig, port: int) -> List[str]:
|
||||
if config.app_type == "wsgi":
|
||||
wrapper_app = self._create_wsgi_wrapper_path(config)
|
||||
app_path = wrapper_app
|
||||
else:
|
||||
app_path = config.app_path
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"uvicorn",
|
||||
app_path,
|
||||
"--host",
|
||||
config.host,
|
||||
"--port",
|
||||
str(port),
|
||||
"--workers",
|
||||
str(config.workers),
|
||||
"--log-level",
|
||||
"warning",
|
||||
"--no-access-log",
|
||||
]
|
||||
|
||||
if config.factory and config.app_type != "wsgi":
|
||||
cmd.append("--factory")
|
||||
|
||||
return cmd
|
||||
|
||||
def _create_wsgi_wrapper_path(self, config: ProcessConfig) -> str:
|
||||
"""
|
||||
Since uvicorn can't directly run WSGI apps, we create a wrapper
|
||||
that imports the WSGI app and wraps it with a2wsgi.
|
||||
"""
|
||||
# For WSGI apps, we'll use a special wrapper module
|
||||
# The wrapper is: pyserve._wsgi_wrapper:create_app
|
||||
# It will be called with app_path as environment variable
|
||||
return "pyserve._wsgi_wrapper:app"
|
||||
|
||||
def get_metrics(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"managed_processes": len(self._processes),
|
||||
"running_processes": sum(1 for p in self._processes.values() if p.is_running),
|
||||
"processes": {name: info.to_dict() for name, info in self._processes.items()},
|
||||
}
|
||||
|
||||
|
||||
_process_manager: Optional[ProcessManager] = None
|
||||
|
||||
|
||||
def get_process_manager() -> ProcessManager:
|
||||
global _process_manager
|
||||
if _process_manager is None:
|
||||
_process_manager = ProcessManager()
|
||||
return _process_manager
|
||||
|
||||
|
||||
async def init_process_manager(
|
||||
port_range: tuple[int, int] = (9000, 9999),
|
||||
health_check_enabled: bool = True,
|
||||
) -> ProcessManager:
|
||||
global _process_manager
|
||||
_process_manager = ProcessManager(
|
||||
port_range=port_range,
|
||||
health_check_enabled=health_check_enabled,
|
||||
)
|
||||
await _process_manager.start()
|
||||
return _process_manager
|
||||
|
||||
|
||||
async def shutdown_process_manager() -> None:
|
||||
global _process_manager
|
||||
if _process_manager:
|
||||
await _process_manager.stop()
|
||||
_process_manager = None
|
||||
+25
-1
@@ -119,6 +119,7 @@ class PyServeServer:
|
||||
self.config = config
|
||||
self.extension_manager = ExtensionManager()
|
||||
self.app: Optional[Starlette] = None
|
||||
self._async_extensions_loaded = False
|
||||
self._setup_logging()
|
||||
self._load_extensions()
|
||||
self._create_app()
|
||||
@@ -133,16 +134,39 @@ class PyServeServer:
|
||||
if ext_config.type == "routing":
|
||||
config.setdefault("default_proxy_timeout", self.config.server.proxy_timeout)
|
||||
|
||||
if ext_config.type == "process_orchestration":
|
||||
continue
|
||||
|
||||
self.extension_manager.load_extension(ext_config.type, config)
|
||||
|
||||
async def _load_async_extensions(self) -> None:
|
||||
if self._async_extensions_loaded:
|
||||
return
|
||||
|
||||
for ext_config in self.config.extensions:
|
||||
if ext_config.type == "process_orchestration":
|
||||
config = ext_config.config.copy()
|
||||
await self.extension_manager.load_extension_async(ext_config.type, config)
|
||||
|
||||
self._async_extensions_loaded = True
|
||||
|
||||
def _create_app(self) -> None:
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncIterator
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: Starlette) -> AsyncIterator[None]:
|
||||
await self._load_async_extensions()
|
||||
logger.info("Async extensions loaded")
|
||||
yield
|
||||
|
||||
routes = [
|
||||
Route("/health", self._health_check, methods=["GET"]),
|
||||
Route("/metrics", self._metrics, methods=["GET"]),
|
||||
Route("/{path:path}", self._catch_all, methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"]),
|
||||
]
|
||||
|
||||
self.app = Starlette(routes=routes)
|
||||
self.app = Starlette(routes=routes, lifespan=lifespan)
|
||||
self.app.add_middleware(PyServeMiddleware, extension_manager=self.extension_manager)
|
||||
|
||||
async def _health_check(self, request: Request) -> Response:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user