From 016dc3d71a1f38437d8d4467c600c7c47202e104 Mon Sep 17 00:00:00 2001 From: Richard Palethorpe Date: Wed, 30 Sep 2026 16:07:36 +0100 Subject: [PATCH] feat(motion): add GEM-X backend and duplex streaming API Add the GEM-X native backend with live pose and root displacement output, bounded duplex WebSocket streaming, session authorization, metrics and usage accounting. Share the embedded bidirectional RPC channel adapter across streaming backends. Include backend build and gallery integration, API documentation, schemas and tests. Assisted-by: Codex:gpt-6 Signed-off-by: Richard Palethorpe --- .dockerignore | 2 + .github/backend-matrix.yml | 61 +++ .github/workflows/bump_deps.yaml | 4 + .gitignore | 5 + Dockerfile | 1 + Makefile | 14 +- README.md | 1 + backend/backend.proto | 14 + backend/go/gemxcpp/CMakeLists.txt | 8 + backend/go/gemxcpp/Makefile | 54 ++ backend/go/gemxcpp/README.md | 19 + backend/go/gemxcpp/gemx.go | 419 +++++++++++++++ backend/go/gemxcpp/gemx_test.go | 160 ++++++ backend/go/gemxcpp/image_projection.go | 69 +++ backend/go/gemxcpp/image_projection_test.go | 84 +++ backend/go/gemxcpp/main.go | 22 + backend/go/gemxcpp/native.go | 52 ++ backend/go/gemxcpp/native/bridge.cpp | 13 + backend/go/gemxcpp/package.sh | 39 ++ backend/go/gemxcpp/package_test.go | 58 +++ backend/go/gemxcpp/root_displacement_test.go | 114 ++++ backend/go/gemxcpp/run.sh | 19 + backend/index.yaml | 62 +++ core/application/application.go | 48 +- core/backend/motion.go | 33 ++ core/config/backend_capabilities.go | 4 + core/config/backend_capabilities_test.go | 14 + core/config/model_config.go | 7 +- core/gallery/importers/gemxcpp.go | 70 +++ core/gallery/importers/gemxcpp_test.go | 45 ++ core/gallery/importers/importers.go | 1 + core/http/app.go | 6 +- core/http/auth/csrf.go | 13 +- core/http/auth/features.go | 6 + core/http/auth/features_motion_test.go | 21 + core/http/auth/header_auth_test.go | 62 +++ core/http/auth/middleware.go | 22 + core/http/auth/motion_permissions_test.go | 86 +++ core/http/auth/permissions.go | 3 +- core/http/auth/websocket_tickets.go | 200 +++++++ .../auth/websocket_tickets_internal_test.go | 88 ++++ .../endpoints/localai/api_instructions.go | 1 + .../localai/api_instructions_test.go | 3 +- core/http/endpoints/localai/motion.go | 414 +++++++++++++++ core/http/endpoints/localai/motion_duplex.go | 357 +++++++++++++ .../endpoints/localai/motion_duplex_test.go | 246 +++++++++ .../endpoints/localai/motion_internal_test.go | 323 ++++++++++++ core/http/middleware/csrf_test.go | 51 ++ core/http/middleware/trace.go | 15 +- core/http/middleware/trace_redact_test.go | 6 + .../http/react-ui/e2e/discover-height.spec.js | 8 + .../public/locales/en/importModel.json | 1 + .../react-ui/public/locales/en/models.json | 4 +- core/http/react-ui/src/pages/ImportModel.jsx | 2 +- core/http/react-ui/src/pages/Models.jsx | 1 + core/http/react-ui/src/utils/capabilities.js | 2 + core/http/route_coverage_test.go | 13 + core/http/routes/localai.go | 7 + core/http/routes/ui_api.go | 1 + core/services/nodes/health_mock_test.go | 4 + core/services/nodes/inflight_test.go | 4 + core/services/worker/free_timeout_test.go | 5 +- core/services/worker/model_stop_test.go | 5 +- core/trace/backend_trace.go | 1 + docs/content/features/backends.md | 4 + docs/content/features/motion.md | 362 +++++++++++++ docs/content/features/runtime-settings.md | 4 +- gallery/index.yaml | 33 ++ pkg/grpc/backend.go | 1 + pkg/grpc/base/base.go | 4 + pkg/grpc/channel_stream.go | 129 +++++ pkg/grpc/channel_stream_test.go | 110 ++++ pkg/grpc/embed.go | 493 +----------------- pkg/grpc/interface.go | 1 + pkg/grpc/motion.go | 101 ++++ pkg/grpc/motion_embed.go | 14 + pkg/grpc/motion_test.go | 66 +++ pkg/model/remote_shutdown_test.go | 7 +- pkg/motion/frame_queue.go | 113 ++++ pkg/motion/frame_queue_test.go | 73 +++ pkg/motion/proto/motion.proto | 75 +++ pkg/motion/validation.go | 36 ++ scripts/lib/backend-filter.mjs | 6 + scripts/lib/backend-filter_test.mjs | 16 + swagger/docs.go | 202 +++++++ swagger/swagger.json | 202 +++++++ swagger/swagger.yaml | 130 +++++ 87 files changed, 5156 insertions(+), 523 deletions(-) create mode 100644 backend/go/gemxcpp/CMakeLists.txt create mode 100644 backend/go/gemxcpp/Makefile create mode 100644 backend/go/gemxcpp/README.md create mode 100644 backend/go/gemxcpp/gemx.go create mode 100644 backend/go/gemxcpp/gemx_test.go create mode 100644 backend/go/gemxcpp/image_projection.go create mode 100644 backend/go/gemxcpp/image_projection_test.go create mode 100644 backend/go/gemxcpp/main.go create mode 100644 backend/go/gemxcpp/native.go create mode 100644 backend/go/gemxcpp/native/bridge.cpp create mode 100755 backend/go/gemxcpp/package.sh create mode 100644 backend/go/gemxcpp/package_test.go create mode 100644 backend/go/gemxcpp/root_displacement_test.go create mode 100755 backend/go/gemxcpp/run.sh create mode 100644 core/backend/motion.go create mode 100644 core/gallery/importers/gemxcpp.go create mode 100644 core/gallery/importers/gemxcpp_test.go create mode 100644 core/http/auth/features_motion_test.go create mode 100644 core/http/auth/header_auth_test.go create mode 100644 core/http/auth/motion_permissions_test.go create mode 100644 core/http/auth/websocket_tickets.go create mode 100644 core/http/auth/websocket_tickets_internal_test.go create mode 100644 core/http/endpoints/localai/motion.go create mode 100644 core/http/endpoints/localai/motion_duplex.go create mode 100644 core/http/endpoints/localai/motion_duplex_test.go create mode 100644 core/http/endpoints/localai/motion_internal_test.go create mode 100644 core/http/middleware/csrf_test.go create mode 100644 docs/content/features/motion.md create mode 100644 pkg/grpc/channel_stream.go create mode 100644 pkg/grpc/channel_stream_test.go create mode 100644 pkg/grpc/motion.go create mode 100644 pkg/grpc/motion_embed.go create mode 100644 pkg/grpc/motion_test.go create mode 100644 pkg/motion/frame_queue.go create mode 100644 pkg/motion/frame_queue_test.go create mode 100644 pkg/motion/proto/motion.proto create mode 100644 pkg/motion/validation.go diff --git a/.dockerignore b/.dockerignore index 5a71590bd1ac..b7ce84903b09 100644 --- a/.dockerignore +++ b/.dockerignore @@ -10,6 +10,8 @@ backend/go/image/stablediffusion-ggml/build/ backend/go/*/build backend/go/kimodocpp/build-* backend/go/kimodocpp/kimodocpp +backend/go/gemxcpp/build-* +backend/go/gemxcpp/gemxcpp backend/go/*/.cache backend/go/*/sources backend/go/*/package diff --git a/.github/backend-matrix.yml b/.github/backend-matrix.yml index 8649209cf706..726907803263 100644 --- a/.github/backend-matrix.yml +++ b/.github/backend-matrix.yml @@ -3616,6 +3616,63 @@ include: dockerfile: "./backend/Dockerfile.golang" context: "./" ubuntu-version: '2404' + # gem-x.cpp: CPU and Vulkan, with native runners for each architecture. + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/amd64' + platform-tag: 'amd64' + tag-latest: 'auto' + tag-suffix: '-cpu-gemxcpp' + runs-on: 'ubuntu-latest' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "gemxcpp" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/arm64' + platform-tag: 'arm64' + tag-latest: 'auto' + tag-suffix: '-cpu-gemxcpp' + runs-on: 'ubuntu-24.04-arm' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "gemxcpp" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' + - build-type: 'vulkan' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/amd64' + platform-tag: 'amd64' + tag-latest: 'auto' + tag-suffix: '-gpu-vulkan-gemxcpp' + runs-on: 'ubuntu-latest' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "gemxcpp" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' + - build-type: 'vulkan' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/arm64' + platform-tag: 'arm64' + tag-latest: 'auto' + tag-suffix: '-gpu-vulkan-gemxcpp' + runs-on: 'ubuntu-24.04-arm' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "gemxcpp" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' # trellis2cpp - build-type: '' cuda-major-version: "" @@ -6551,6 +6608,10 @@ includeDarwin: tag-suffix: "-cpu-darwin-arm64-kimodocpp" build-type: "cpu" lang: "go" + - backend: "gemxcpp" + tag-suffix: "-cpu-darwin-arm64-gemxcpp" + build-type: "cpu" + lang: "go" - backend: "diffusers" tag-suffix: "-metal-darwin-arm64-diffusers" build-type: "mps" diff --git a/.github/workflows/bump_deps.yaml b/.github/workflows/bump_deps.yaml index e8569473a2be..de3a8b74d515 100644 --- a/.github/workflows/bump_deps.yaml +++ b/.github/workflows/bump_deps.yaml @@ -94,6 +94,10 @@ jobs: variable: "KIMODO_VERSION" branch: "main" file: "backend/go/kimodocpp/Makefile" + - repository: "localai-org/gem-x.cpp" + variable: "GEMX_VERSION" + branch: "main" + file: "backend/go/gemxcpp/Makefile" - repository: "mudler/go-piper" variable: "PIPER_VERSION" branch: "master" diff --git a/.gitignore b/.gitignore index f59bbe45337f..e0a18e1ab820 100644 --- a/.gitignore +++ b/.gitignore @@ -134,3 +134,8 @@ formal-verification/out/ # root, which is what a contributor testing a build does. Nothing under here is # source: it is the instance's own models, outputs, traces and identity. /data/ + +/backend/go/gemxcpp/sources/ +/backend/go/gemxcpp/build-*/ +/backend/go/gemxcpp/package/ +/backend/go/gemxcpp/gemxcpp diff --git a/Dockerfile b/Dockerfile index 4ca9b32791ab..4d49e38a11e0 100644 --- a/Dockerfile +++ b/Dockerfile @@ -346,6 +346,7 @@ COPY ./.git ./.git # Some of the Go backends use libs from the main src, we could further optimize the caching by building the CPP backends before here COPY ./pkg/grpc ./pkg/grpc +COPY ./pkg/motion ./pkg/motion COPY ./pkg/utils ./pkg/utils RUN ls -l ./ diff --git a/Makefile b/Makefile index 7c05fda6dd3d..79c89719fbab 100644 --- a/Makefile +++ b/Makefile @@ -591,6 +591,7 @@ protogen-go: protoc install-go-tools # shell-profile change. PATH="$$(go env GOPATH)/bin:$$PATH" ./protoc --experimental_allow_proto3_optional -Ibackend/ --go_out=pkg/grpc/proto/ --go_opt=paths=source_relative --go-grpc_out=pkg/grpc/proto/ --go-grpc_opt=paths=source_relative \ backend/backend.proto + PATH="$$(go env GOPATH)/bin:$$PATH" ./protoc -Ipkg/motion/proto --go_out=pkg/motion/proto --go_opt=paths=source_relative pkg/motion/proto/motion.proto core/config/inference_defaults.json: ## Fetch inference defaults from unsloth (only if missing) $(GOCMD) generate ./core/config/... @@ -604,7 +605,7 @@ generate-force: ## Re-fetch inference defaults from unsloth (always) .PHONY: protogen-go-clean protogen-go-clean: - $(RM) pkg/grpc/proto/backend.pb.go pkg/grpc/proto/backend_grpc.pb.go + $(RM) pkg/grpc/proto/backend.pb.go pkg/grpc/proto/backend_grpc.pb.go pkg/motion/proto/motion.pb.go $(RM) bin/* prepare-test-extra: protogen-python @@ -641,6 +642,7 @@ prepare-test-extra: protogen-python $(MAKE) -C backend/go/locate-anything-cpp $(MAKE) -C backend/go/trellis2cpp $(MAKE) -C backend/go/kimodocpp + $(MAKE) -C backend/go/gemxcpp $(MAKE) -C backend/go/valkey-store test-extra: prepare-test-extra @@ -680,6 +682,7 @@ test-extra: prepare-test-extra $(MAKE) -C backend/go/nemo-speech-cpp test $(MAKE) -C backend/go/trellis2cpp test $(MAKE) -C backend/go/kimodocpp test + $(MAKE) -C backend/go/gemxcpp test $(MAKE) -C backend/go/valkey-store test ## @@ -1336,6 +1339,7 @@ BACKEND_HUGGINGFACE = huggingface|golang|.|false|true BACKEND_SILERO_VAD = silero-vad|golang|.|false|true BACKEND_STABLEDIFFUSION_GGML = stablediffusion-ggml|golang|.|--progress=plain|true BACKEND_TRELLIS2CPP = trellis2cpp|golang|.|--progress=plain|true +BACKEND_GEMXCPP = gemxcpp|golang|.|--progress=plain|true BACKEND_KIMODOCPP = kimodocpp|golang|.|--progress=plain|true BACKEND_WHISPER = whisper|golang|.|false|true BACKEND_CRISPASR = crispasr|golang|.|false|true @@ -1444,6 +1448,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_SILERO_VAD))) $(eval $(call generate-docker-build-target,$(BACKEND_STABLEDIFFUSION_GGML))) $(eval $(call generate-docker-build-target,$(BACKEND_TRELLIS2CPP))) $(eval $(call generate-docker-build-target,$(BACKEND_KIMODOCPP))) +$(eval $(call generate-docker-build-target,$(BACKEND_GEMXCPP))) .NOTPARALLEL: backends/kimodocpp backends/kimodocpp-darwin docker-build-backends: docker-build-kimodocpp @@ -1451,6 +1456,13 @@ backends/kimodocpp-darwin: BACKEND=kimodocpp BUILD_TYPE=cpu $(MAKE) build-darwin-go-backend ./local-ai backends install "ocifile://$(abspath ./backend-images/kimodocpp.tar)" +.NOTPARALLEL: backends/gemxcpp backends/gemxcpp-darwin +docker-build-backends: docker-build-gemxcpp + +backends/gemxcpp-darwin: + BACKEND=gemxcpp BUILD_TYPE=cpu $(MAKE) build-darwin-go-backend + ./local-ai backends install "ocifile://$(abspath ./backend-images/gemxcpp.tar)" + $(eval $(call generate-docker-build-target,$(BACKEND_WHISPER))) $(eval $(call generate-docker-build-target,$(BACKEND_CRISPASR))) $(eval $(call generate-docker-build-target,$(BACKEND_PARAKEET_CPP))) diff --git a/README.md b/README.md index f1c6c62900ab..7548f2c4fa83 100644 --- a/README.md +++ b/README.md @@ -256,6 +256,7 @@ Most backends wrap a best-in-class upstream engine. A handful of them are native | [face-detect.cpp](https://github.com/mudler/face-detect.cpp) | Face detection, recognition, demographics and anti-spoofing (SCRFD/ArcFace, YuNet/SFace), replacing the Python insightface backend | | [free-splatter.cpp](https://github.com/localai-org/free-splatter.cpp) | Pose-free 3D reconstruction (FreeSplatter): turns a handful of plain photos into 3D Gaussians, no camera poses or GPU required | | [trellis2.cpp](https://github.com/localai-org/trellis2cpp) | C++/GGML port of Microsoft TRELLIS.2: single-image to textured 3D mesh (GLB with PBR materials) | +| [gem-x.cpp](https://github.com/localai-org/gem-x.cpp) | Live human motion capture with SOMA and SMPL pose streams (CPU/Vulkan) | | [kimodo.cpp](https://github.com/localai-org/kimodo.cpp) | C++/GGML text-to-motion on CPU and Vulkan, exported as animated skeleton GLB | | [privacy-filter.cpp](https://github.com/localai-org/privacy-filter.cpp) | Standalone GGML PII/NER token-classification engine powering LocalAI's PII redaction tier | | [LocalVQE](https://github.com/localai-org/LocalVQE) | Joint acoustic echo cancellation, noise suppression, and dereverberation | diff --git a/backend/backend.proto b/backend/backend.proto index 614f006e9fe1..aae0b4205901 100644 --- a/backend/backend.proto +++ b/backend/backend.proto @@ -7,6 +7,7 @@ option java_outer_classname = "LocalAIBackend"; package backend; + service Backend { rpc Health(HealthMessage) returns (Reply) {} rpc Free(HealthMessage) returns (Result) {} @@ -18,6 +19,7 @@ service Backend { rpc UpscaleImage(UpscaleImageRequest) returns (Result) {} rpc GenerateVideo(GenerateVideoRequest) returns (Result) {} rpc Generate3D(Generate3DRequest) returns (Result) {} + rpc MotionStream(stream MotionRequest) returns (stream MotionResponse) {} rpc Animate3D(Animate3DRequest) returns (Result) {} rpc AudioTranscription(TranscriptRequest) returns (TranscriptResult) {} rpc AudioTranscriptionStream(TranscriptRequest) returns (stream TranscriptStreamResponse) {} @@ -1574,3 +1576,15 @@ message ForwardReply { repeated ForwardHeader headers = 2; bytes body_chunk = 3; } + +// Configuration is first; later messages contain input only. +message MotionRequest { + string model_identity = 1; + string profile = 2; + bytes input = 3; // Serialized public localai.motion.v1.Input. +} +message MotionResponse { + bytes output = 1; // Serialized public localai.motion.v1.Output. + // Backend-only bounded stage timings; not part of the public wire protocol. + map stage_ms = 2; +} diff --git a/backend/go/gemxcpp/CMakeLists.txt b/backend/go/gemxcpp/CMakeLists.txt new file mode 100644 index 000000000000..725603177e0c --- /dev/null +++ b/backend/go/gemxcpp/CMakeLists.txt @@ -0,0 +1,8 @@ +# SPDX-License-Identifier: MIT +cmake_minimum_required(VERSION 3.24) +project(localai_gemx LANGUAGES C CXX) +set(GEMX_SANITIZERS OFF CACHE BOOL "" FORCE) +set(BUILD_TESTING OFF CACHE BOOL "" FORCE) +set(GGML_METAL OFF CACHE BOOL "" FORCE) +add_subdirectory(sources/gem-x.cpp) +target_sources(gemx PRIVATE native/bridge.cpp) diff --git a/backend/go/gemxcpp/Makefile b/backend/go/gemxcpp/Makefile new file mode 100644 index 000000000000..7d2599fcce7d --- /dev/null +++ b/backend/go/gemxcpp/Makefile @@ -0,0 +1,54 @@ +GEMX_REPO?=https://github.com/localai-org/gem-x.cpp +GEMX_VERSION?=576cab62d18829ae00df5959b8ccf0a1032e0c87 +BUILD_TYPE?= +GOCMD?=go +GO_TAGS?= +JOBS?=4 +CMAKE_ARGS?= +BUILD_DIR?=build-$(if $(filter vulkan,$(BUILD_TYPE)),vulkan,cpu) + +CMAKE_ARGS+=-DCMAKE_BUILD_TYPE=Release -DGGML_NATIVE=OFF -DGGML_METAL=OFF +ifeq ($(shell uname -s),Darwin) +CMAKE_ARGS+=-DCMAKE_BUILD_WITH_INSTALL_RPATH=ON -DCMAKE_INSTALL_RPATH=@loader_path +endif +ifeq ($(BUILD_TYPE),vulkan) +CMAKE_ARGS+=-DGEMX_VULKAN=ON +else +CMAKE_ARGS+=-DGEMX_VULKAN=OFF +endif + +.PHONY: all build package test clean purge FORCE +all: build + +FORCE: + +sources/gem-x.cpp/.checkout-ready: FORCE + @test -d sources/gem-x.cpp/.git || git clone $(GEMX_REPO) sources/gem-x.cpp + @if test -f $@ && test "$$(git -C sources/gem-x.cpp rev-parse HEAD)" = "$(GEMX_VERSION)"; then exit 0; fi; \ + cd sources/gem-x.cpp && (git diff --quiet && git diff --cached --quiet || \ + { echo "gem-x.cpp has local source changes; preserve them and run make clean before updating the pin" >&2; exit 1; }) && \ + git fetch origin $(GEMX_VERSION) && git checkout $(GEMX_VERSION) && \ + git submodule update --init --recursive --depth 1 && touch .checkout-ready + +$(BUILD_DIR)/.built: sources/gem-x.cpp/.checkout-ready CMakeLists.txt native/bridge.cpp Makefile + cmake -S . -B $(BUILD_DIR) $(CMAKE_ARGS) + cmake --build $(BUILD_DIR) -j$(JOBS) + touch $@ + + +gemxcpp: main.go gemx.go native.go image_projection.go $(BUILD_DIR)/.built + CGO_ENABLED=0 $(GOCMD) build -tags "$(GO_TAGS)" -o $@ ./ + +package: gemxcpp + BUILD_DIR=$(BUILD_DIR) bash package.sh + +build: package + +test: + CGO_ENABLED=0 $(GOCMD) test -v ./... + +purge: + rm -rf build-native build-cpu build-vulkan build-native-avx2 build-cpu-avx2 build-vulkan-avx2 + +clean: purge + rm -rf package gemxcpp sources diff --git a/backend/go/gemxcpp/README.md b/backend/go/gemxcpp/README.md new file mode 100644 index 000000000000..0f2e4f785cdf --- /dev/null +++ b/backend/go/gemxcpp/README.md @@ -0,0 +1,19 @@ +# GEM-X backend + +Resident native `gemx_live_*` inference via PureGo. See +[Motion Capture](../../../docs/content/features/motion.md) for configuration and +protocol. The upstream pin lives in `Makefile` and is registered with the daily +bump workflow. Linux CPU/Vulkan and Darwin CPU packages are built; Metal is not +supported by the upstream runtime. Model assets are not needed for unit tests. + +```sh +make -C backend/go/gemxcpp test +make backends/gemxcpp +``` + +The adapter owns one native pipeline per loaded model, with one active session. +Each instance loads GEM, ViTPose and YOLOX. No SAM3D Body, mesh export, SONIC, +physics, or Python dependency is introduced. The C library has no immediate +inference cancellation; cancellation discards results and releases the session +when the native call returns. The public stream schema lives in +`pkg/motion/proto/motion.proto`, independently of backend RPC envelopes. diff --git a/backend/go/gemxcpp/gemx.go b/backend/go/gemxcpp/gemx.go new file mode 100644 index 000000000000..6eac877674b8 --- /dev/null +++ b/backend/go/gemxcpp/gemx.go @@ -0,0 +1,419 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "context" + "encoding/json" + "fmt" + "io" + "math" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "sync" + + "github.com/mudler/LocalAI/pkg/grpc/base" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + motionutil "github.com/mudler/LocalAI/pkg/motion" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/proto" +) + +type GemX struct { + base.Base + mu sync.Mutex + pipeline uintptr + active bool + definition map[string]json.RawMessage +} + +type loadOptions struct { + gem, pose, detector, module, device string + threads, window, cadence, selection, precision, index uint32 + gap int64 +} + +func parseOptions(o *pb.ModelOptions) (loadOptions, error) { + c := loadOptions{gem: o.ModelFile, module: os.Getenv("GEMX_MODULE"), device: "CPU", threads: uint32(min(max(int(o.Threads), 1), 8)), window: 30, cadence: 1, selection: 1, gap: 2000000} + if os.Getenv("GEMX_DEVICE") != "" { + c.device = os.Getenv("GEMX_DEVICE") + } + for _, v := range o.Options { + k, v, ok := strings.Cut(v, ":") + if !ok { + return c, fmt.Errorf("invalid GEM-X option") + } + switch k { + case "vitpose": + c.pose = v + case "yolox": + c.detector = v + case "device": + switch v { + case "cpu": + c.device = "CPU" + case "vulkan": + c.device = "Vulkan" + default: + return c, fmt.Errorf("device must be cpu or vulkan") + } + case "selection": + switch v { + case "continuity": + c.selection = 1 + case "parity": + c.selection = 0 + default: + return c, fmt.Errorf("selection must be continuity or parity") + } + case "precision": + switch v { + case "strict": + c.precision = 0 + case "backend_default": + c.precision = 1 + default: + return c, fmt.Errorf("precision must be strict or backend_default") + } + case "window", "detector_interval", "device_index", "max_gap_us": + n, err := strconv.ParseUint(v, 10, 32) + if err != nil { + return c, fmt.Errorf("invalid %s", k) + } + switch k { + case "window": + c.window = uint32(n) + case "detector_interval": + c.cadence = uint32(n) + case "device_index": + c.index = uint32(n) + case "max_gap_us": + c.gap = int64(n) + } + default: + return c, fmt.Errorf("unsupported GEM-X option %q", k) + } + } + if c.window < 2 || c.window > 120 || c.cadence < 1 || c.cadence > 30 || c.gap <= 0 { + return c, fmt.Errorf("window must be 2..120, detector_interval 1..30, max_gap_us positive") + } + for _, p := range []*string{&c.gem, &c.pose, &c.detector} { + if *p == "" { + return c, fmt.Errorf("model, vitpose and yolox are required") + } + if !filepath.IsAbs(*p) { + *p = filepath.Join(o.ModelPath, *p) + } + info, err := os.Stat(*p) + if err != nil || !info.Mode().IsRegular() { + return c, fmt.Errorf("model asset is not a regular file: %s", *p) + } + } + if c.module == "" { + return c, fmt.Errorf("GEMX_MODULE is required") + } + return c, nil +} +func (g *GemX) Load(o *pb.ModelOptions) error { + g.mu.Lock() + defer g.mu.Unlock() + if g.active { + return status.Error(codes.ResourceExhausted, "GEM-X session is active") + } + c, err := parseOptions(o) + if err != nil { + return err + } + if g.pipeline != 0 { + return fmt.Errorf("GEM-X model already loaded") + } + buf := make([]byte, 2048) + scalars := []uint32{c.index, c.threads, c.window, c.cadence, c.selection, c.precision} + code := liveCreate(c.gem, c.pose, c.detector, c.module, c.device, &scalars[0], c.gap, &g.pipeline, &buf[0], uint64(len(buf))) + if err := nativeError(code, buf); err != nil { + return err + } + var n uint64 + if err := nativeError(liveDefinition(g.pipeline, nil, 0, &n, &buf[0], uint64(len(buf))), buf); err != nil { + g.free() + return err + } + if n < 2 || n > 1<<20 { + g.free() + return fmt.Errorf("invalid native definition size") + } + def := make([]byte, n) + if err := nativeError(liveDefinition(g.pipeline, &def[0], n, &n, &buf[0], uint64(len(buf))), buf); err != nil { + g.free() + return err + } + if err := json.Unmarshal(def[:len(def)-1], &g.definition); err != nil { + g.free() + return err + } + xlog.Info("GEM-X loaded", "device", c.device, "window", c.window, "detector_interval", c.cadence) + return nil +} +func (g *GemX) Busy() bool { g.mu.Lock(); defer g.mu.Unlock(); return g.active } + +func (g *GemX) free() { + if g.pipeline != 0 { + liveDestroy(g.pipeline) + g.pipeline = 0 + } +} +func (g *GemX) Free() error { + g.mu.Lock() + defer g.mu.Unlock() + if g.active { + return status.Error(codes.ResourceExhausted, "GEM-X session is active") + } + g.free() + return nil +} + +func supportsRootDisplacement(raw map[string]json.RawMessage, profile string) bool { + var channels map[string]int + return resultIntervalStart != nil && profile == "smpl24" && json.Unmarshal(raw["channels"], &channels) == nil && channels["root_displacement"] == 3 +} + +func publicDefinition(raw map[string]json.RawMessage, profile string) (*motion.Definition, error) { + var fields map[string]json.RawMessage + if err := json.Unmarshal(raw[profile], &fields); err != nil { + return nil, err + } + d := &motion.Definition{Profile: profile, Conventions: map[string]string{"time_unit": "microseconds", "time_origin": "caller-declared", "temporal_policy": "accepted-frame-index", "continuous_world_trajectory": "false"}} + for key, dst := range map[string]any{"schema": &d.Schema, "joint_names": &d.JointNames, "parents": &d.Parents, "root": &d.Root, "rest_local_translations": &d.RestLocalTranslations, "rest_local_rotations": &d.RestLocalRotations} { + if err := json.Unmarshal(fields[key], dst); err != nil { + return nil, fmt.Errorf("native definition %s: %w", key, err) + } + } + for _, k := range []string{"units", "handedness", "basis", "space", "quaternion_order", "quaternion_policy", "shape_policy", "reference_shape", "anchor_basis", "reconstruction", "rest_rotation_policy"} { + if val := fields[k]; val != nil { + var s string + if err := json.Unmarshal(val, &s); err != nil { + return nil, err + } + d.Conventions[k] = s + } + } + d.Channels = []string{"positions", "root_translation", "box", "image_positions"} + if supportsRootDisplacement(raw, profile) { + d.Channels = append(d.Channels, "root_displacement") + d.Conventions["root_displacement"] = "metres, end-anchor-local, closed source interval, no cross-epoch continuity" + } + d.Conventions["image_positions"] = "source-pixel xy, joint_names order, top-left origin, unmirrored" + if profile == "soma77" { + d.Channels = append(d.Channels, "local_rotations", "local_translations", "root_axis_angle") + } else { + d.Channels = append(d.Channels, "anchor") + } + if len(d.JointNames) != len(d.Parents) || len(d.RestLocalTranslations) != len(d.Parents)*3 || len(d.RestLocalRotations) != len(d.Parents)*4 { + return nil, fmt.Errorf("invalid native topology") + } + return d, nil +} + +func (g *GemX) MotionStream(ctx context.Context, recv func() (*pb.MotionRequest, error), send func(*pb.MotionResponse) error) error { + first, err := recv() + if err != nil { + return err + } + if first.Profile != "soma77" && first.Profile != "smpl24" { + return status.Error(codes.InvalidArgument, "profile must be soma77 or smpl24") + } + if len(first.Input) != 0 { + return status.Error(codes.InvalidArgument, "configuration must precede frames") + } + g.mu.Lock() + if g.active || g.pipeline == 0 { + g.mu.Unlock() + return status.Error(codes.ResourceExhausted, "GEM-X is unavailable or already streaming") + } + g.active = true + g.mu.Unlock() + defer func() { g.mu.Lock(); g.active = false; g.mu.Unlock() }() + buf := make([]byte, 2048) + if err := nativeError(liveReset(g.pipeline, &buf[0], uint64(len(buf))), buf); err != nil { + return err + } + def, err := publicDefinition(g.definition, first.Profile) + if err != nil { + return err + } + projection, err := projectionForProfile(g.definition, first.Profile) + if err != nil { + return err + } + + emit := func(o *motion.Output, metrics map[string]float64) error { + b, err := proto.Marshal(o) + if err != nil { + return err + } + return send(&pb.MotionResponse{Output: b, StageMs: metrics}) + } + if err := emit(&motion.Output{Payload: &motion.Output_Definition{Definition: def}}, nil); err != nil { + return err + } + displacement := supportsRootDisplacement(g.definition, first.Profile) + for { + if err := ctx.Err(); err != nil { + return err + } + req, err := recv() + if err == io.EOF { + return nil + } + if err != nil { + return err + } + if req.Profile != "" { + return status.Error(codes.InvalidArgument, "configuration cannot change") + } + input := &motion.Input{} + if err := proto.Unmarshal(req.Input, input); err != nil { + return status.Error(codes.InvalidArgument, "invalid motion input") + } + if input.GetResetState() { + if err := nativeError(liveReset(g.pipeline, &buf[0], uint64(len(buf))), buf); err != nil { + return err + } + if err := emit(&motion.Output{Payload: &motion.Output_Event{Event: &motion.Event{Type: "reset"}}}, nil); err != nil { + return err + } + continue + } + frame := input.GetFrame() + if err := validateFrame(frame); err != nil { + return status.Error(codes.InvalidArgument, err.Error()) + } + var box *float32 + if len(frame.SubjectBox) > 0 { + box = &frame.SubjectBox[0] + } + var result uintptr + code := liveSubmit(g.pipeline, &frame.Rgb[0], uint64(len(frame.Rgb)), frame.Width, frame.Height, uint64(frame.Width)*3, frame.Sequence, frame.SourceTimeUs, box, uint64(len(frame.SubjectBox)), frame.SubjectId, &result, &buf[0], uint64(len(buf))) + runtime.KeepAlive(frame) + if err := nativeError(code, buf); err != nil { + return err + } + projection.width, projection.height = frame.Width, frame.Height + out, metrics, err := readResult(result, first.Profile, projection, displacement) + resultDestroy(result) + if err != nil { + return err + } + if err := emit(out, metrics); err != nil { + return err + } + } +} +func validateFrame(f *motion.Frame) error { return motionutil.ValidateFrame(f) } + +func readResult(r uintptr, profile string, projection poseProjection, displacement bool) (*motion.Output, map[string]float64, error) { + buf := make([]byte, 2048) + p := &motion.Pose{} + var outcome uint32 + if err := nativeError(resultInfo(r, &p.Sequence, &p.SourceTimeUs, &p.Epoch, &p.TrackEpoch, &outcome, &p.Flags, &buf[0], uint64(len(buf))), buf); err != nil { + return nil, nil, err + } + copyChannel := func(channel uint32, max uint64) ([]float32, error) { + var n uint64 + if err := nativeError(resultCopy(r, channel, nil, 0, &n, &buf[0], uint64(len(buf))), buf); err != nil { + return nil, err + } + if n > max { + return nil, fmt.Errorf("native channel exceeds expected size") + } + if n == 0 { + return nil, nil + } + values := make([]float32, n) + if err := nativeError(resultCopy(r, channel, &values[0], n, &n, &buf[0], uint64(len(buf))), buf); err != nil { + return nil, err + } + for _, v := range values { + if math.IsNaN(float64(v)) || math.IsInf(float64(v), 0) { + return nil, fmt.Errorf("nonfinite native output") + } + } + return values, nil + } + metrics := map[string]float64{} + times, err := copyChannel(11, 9) + if err != nil { + return nil, nil, err + } + if len(times) >= 5 { + for i, k := range []string{"detector", "vitpose", "gem", "world", "camera"} { + metrics[k] = float64(times[i]) + } + } + if outcome != 1 { + names := map[uint32]string{0: "warmup", 2: "lost", 3: "ambiguous"} + name, ok := names[outcome] + if !ok { + return nil, nil, fmt.Errorf("unknown native outcome") + } + return &motion.Output{Payload: &motion.Output_Event{Event: &motion.Event{Type: name, Sequence: p.Sequence, SourceTimeUs: p.SourceTimeUs, Epoch: p.Epoch, TrackEpoch: p.TrackEpoch, Flags: p.Flags}}}, metrics, nil + } + type channel struct { + id uint32 + n uint64 + dst *[]float32 + } + channels := []channel{{6, 3, &p.RootTranslation}, {10, 4, &p.Box}} + if profile == "soma77" { + channels = append(channels, + channel{0, 231, &p.Positions}, + channel{1, 308, &p.LocalRotations}, + channel{2, 231, &p.LocalTranslations}, + channel{5, 3, &p.RootAxisAngle}) + } else { + channels = append(channels, + channel{3, 72, &p.Positions}, + channel{4, 4, &p.Anchor}) + } + for _, c := range channels { + v, err := copyChannel(c.id, c.n) + if err != nil { + return nil, nil, err + } + if uint64(len(v)) != c.n { + return nil, nil, fmt.Errorf("incomplete native pose") + } + *c.dst = v + } + if displacement && p.Flags&1 == 0 { + if err := nativeError(resultIntervalStart(r, &p.DisplacementStartTimeUs, &buf[0], uint64(len(buf))), buf); err != nil { + return nil, nil, err + } + if p.DisplacementStartTimeUs < 0 || p.DisplacementStartTimeUs >= p.SourceTimeUs { + return nil, nil, fmt.Errorf("invalid native displacement interval") + } + delta, err := copyChannel(14, 3) + if err != nil { + return nil, nil, err + } + if len(delta) != 3 { + return nil, nil, fmt.Errorf("incomplete native displacement") + } + p.RootDisplacement = delta + } + camera, err := copyChannel(7, 231) + if err != nil { + return nil, nil, err + } + translation, err := copyChannel(8, 3) + if err != nil { + return nil, nil, err + } + p.ImagePositions = projection.project(camera, translation) + return &motion.Output{Payload: &motion.Output_Pose{Pose: p}}, metrics, nil +} diff --git a/backend/go/gemxcpp/gemx_test.go b/backend/go/gemxcpp/gemx_test.go new file mode 100644 index 000000000000..46f5ce18ef3c --- /dev/null +++ b/backend/go/gemxcpp/gemx_test.go @@ -0,0 +1,160 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "context" + "encoding/json" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/protobuf/proto" + "io" + "math" + "os" + "path/filepath" + "testing" +) + +func TestGemX(t *testing.T) { RegisterFailHandler(Fail); RunSpecs(t, "GEM-X backend") } + +var _ = Describe("GEM-X adapter", func() { + It("rejects malformed RGB and subject coordinates before native inference", func() { + f := &motion.Frame{Width: 8, Height: 8, Rgb: make([]byte, 192)} + Expect(validateFrame(f)).To(Succeed()) + f.Width = math.MaxUint32 + Expect(validateFrame(f)).To(HaveOccurred()) + f.Width = 8 + f.SubjectBox = []float32{0, 0, float32(math.NaN()), 1} + Expect(validateFrame(f)).To(HaveOccurred()) + }) + It("validates all component files and bounded configuration", func() { + dir := GinkgoT().TempDir() + for _, name := range []string{"gem", "pose", "detector"} { + Expect(os.WriteFile(filepath.Join(dir, name), []byte("fixture"), 0600)).To(Succeed()) + } + GinkgoT().Setenv("GEMX_MODULE", dir) + o := &pb.ModelOptions{ModelFile: "gem", ModelPath: dir, Threads: 32, Options: []string{"vitpose:pose", "yolox:detector"}} + c, err := parseOptions(o) + Expect(err).NotTo(HaveOccurred()) + Expect(c.threads).To(Equal(uint32(8))) + Expect(c.selection).To(Equal(uint32(1))) + o.Options = append(o.Options, "window:121") + _, err = parseOptions(o) + Expect(err).To(HaveOccurred()) + o.Options = []string{"vitpose:pose", "yolox:missing"} + _, err = parseOptions(o) + Expect(err).To(HaveOccurred()) + }) + It("publishes only skeleton metadata, never native private configuration", func() { + var raw map[string]json.RawMessage + Expect(json.Unmarshal([]byte(`{"config":{"model_paths":{"gem":"/private/model"}},"smpl24":{"schema":"test.v1","joint_names":["pelvis"],"parents":[-1],"root":0,"rest_local_translations":[0,0,0],"rest_local_rotations":[1,0,0,0],"quaternion_order":"wxyz","space":"root-local"}}`), &raw)).To(Succeed()) + d, err := publicDefinition(raw, "smpl24") + Expect(err).NotTo(HaveOccurred()) + Expect(d.Conventions["quaternion_order"]).To(Equal("wxyz")) + bytes, err := proto.Marshal(d) + Expect(err).NotTo(HaveOccurred()) + Expect(string(bytes)).NotTo(ContainSubstring("private")) + Expect(d.Channels).NotTo(ContainElement("local_rotations")) + }) + It("rejects unsupported profiles before touching native state", func() { + g := &GemX{} + err := g.MotionStream(context.Background(), func() (*pb.MotionRequest, error) { return &pb.MotionRequest{Profile: "robot"}, nil }, nil) + Expect(err).To(HaveOccurred()) + }) + It("releases session ownership after a client closes during warmup", func() { + oldReset := liveReset + DeferCleanup(func() { liveReset = oldReset }) + liveReset = func(uintptr, *byte, uint64) int32 { return 0 } + raw := map[string]json.RawMessage{"soma77": json.RawMessage(`{"schema":"test","joint_names":["root"],"parents":[-1],"root":0,"rest_local_translations":[0,0,0],"rest_local_rotations":[0,0,0,1]}`)} + g := &GemX{pipeline: 1, definition: raw} + calls := 0 + sent := 0 + err := g.MotionStream(context.Background(), func() (*pb.MotionRequest, error) { + calls++ + if calls == 1 { + return &pb.MotionRequest{Profile: "soma77"}, nil + } + return nil, io.EOF + }, func(r *pb.MotionResponse) error { + out := &motion.Output{} + Expect(proto.Unmarshal(r.Output, out)).To(Succeed()) + Expect(out.GetDefinition()).NotTo(BeNil()) + sent++ + return nil + }) + Expect(err).NotTo(HaveOccurred()) + Expect(sent).To(Equal(1)) + Expect(g.active).To(BeFalse()) + }) + It("binds the real native ABI and rejects invalid assets when explicitly enabled", func() { + library := os.Getenv("GEMX_TEST_LIBRARY") + if library == "" { + Skip("set GEMX_TEST_LIBRARY to a built library including the LocalAI bridge") + } + Expect(loadNativeLibrary(library)).To(Succeed()) + var handle uintptr + buf := make([]byte, 2048) + opts := []uint32{0, 1, 30, 1, 1, 0} + code := liveCreate("/missing/gem", "/missing/pose", "/missing/detector", "/missing/module", "CPU", &opts[0], 2000000, &handle, &buf[0], uint64(len(buf))) + Expect(code).NotTo(BeZero()) + Expect(handle).To(BeZero()) + Expect(nativeError(code, buf)).To(HaveOccurred()) + }) + It("streams real model frames with displacement intervals fenced by resets", func() { + models := os.Getenv("GEMX_TEST_MODELS") + if models == "" { + Skip("set GEMX_TEST_MODELS and GEMX_TEST_LIBRARY for native inference") + } + Expect(loadNativeLibrary(os.Getenv("GEMX_TEST_LIBRARY"))).To(Succeed()) + g := &GemX{} + Expect(g.Load(&pb.ModelOptions{ModelFile: "gem-x-contact-f32.gguf", ModelPath: models, Threads: 4, Options: []string{"vitpose:vitpose-f32.gguf", "yolox:yolox-f32.gguf", "device:vulkan"}})).To(Succeed()) + DeferCleanup(func() { Expect(g.Free()).To(Succeed()) }) + requests := []*pb.MotionRequest{{Profile: "smpl24"}} + timestamps := []int64{0, 100000, 350000, 400000, 3000000, 3100000} + for i, timestamp := range timestamps { + f := &motion.Frame{Width: 16, Height: 16, Rgb: make([]byte, 16*16*3), Sequence: uint64(i + 1), SourceTimeUs: timestamp, SubjectBox: []float32{0, 0, 15, 15}, SubjectId: 1} + if i >= 3 { + f.SubjectId = 2 + } + data, err := proto.Marshal(&motion.Input{Payload: &motion.Input_Frame{Frame: f}}) + Expect(err).NotTo(HaveOccurred()) + requests = append(requests, &pb.MotionRequest{Input: data}) + } + var outputs []*motion.Output + Expect(g.MotionStream(context.Background(), func() (*pb.MotionRequest, error) { + if len(requests) == 0 { + return nil, io.EOF + } + r := requests[0] + requests = requests[1:] + return r, nil + }, func(r *pb.MotionResponse) error { + o := &motion.Output{} + Expect(proto.Unmarshal(r.Output, o)).To(Succeed()) + outputs = append(outputs, o) + return nil + })).To(Succeed()) + Expect(outputs).To(HaveLen(7)) + Expect(outputs[0].GetDefinition().JointNames).To(HaveLen(24)) + Expect(outputs[1].GetEvent().Type).To(Equal("warmup")) + Expect(outputs[2].GetPose().Positions).To(HaveLen(72)) + Expect(outputs[2].GetPose().Anchor).To(HaveLen(4)) + Expect(outputs[0].GetDefinition().Conventions["image_positions"]).To(Equal("source-pixel xy, joint_names order, top-left origin, unmirrored")) + Expect(outputs[2].GetPose().ImagePositions).To(HaveLen(48)) + Expect(outputs[2].GetPose().SourceTimeUs).To(Equal(int64(100000))) + Expect(outputs[4].GetEvent().Type).To(Equal("warmup")) // subject change + Expect(outputs[5].GetEvent().Type).To(Equal("warmup")) // source-time gap + Expect(outputs[6].GetPose().Epoch).To(BeNumerically(">", outputs[3].GetPose().Epoch)) + if supportsRootDisplacement(g.definition, "smpl24") { + Expect(outputs[0].GetDefinition().Channels).To(ContainElement("root_displacement")) + for _, index := range []int{2, 3, 6} { + p := outputs[index].GetPose() + Expect(p.RootDisplacement).To(HaveLen(3)) + Expect(p.DisplacementStartTimeUs).To(Equal(timestamps[index-2])) + Expect(p.RootTranslation).To(Equal([]float32{0, 0, 0})) + } + } + }) + +}) diff --git a/backend/go/gemxcpp/image_projection.go b/backend/go/gemxcpp/image_projection.go new file mode 100644 index 000000000000..9898343f62b8 --- /dev/null +++ b/backend/go/gemxcpp/image_projection.go @@ -0,0 +1,69 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "encoding/json" + "fmt" + "math" +) + +type poseProjection struct { + joints []int + width, height uint32 +} + +func projectionForProfile(definition map[string]json.RawMessage, profile string) (poseProjection, error) { + var result poseProjection + + if profile == "soma77" { + result.joints = make([]int, 77) + for i := range result.joints { + result.joints[i] = i + } + return result, nil + } + + var smpl struct { + Mapping []int `json:"mapping"` + } + + if err := json.Unmarshal(definition["smpl24"], &smpl); err != nil { + return result, err + } + if len(smpl.Mapping) != 24 { + return result, fmt.Errorf("invalid SMPL image projection mapping") + } + for _, joint := range smpl.Mapping { + if joint < 0 || joint >= 77 { + return result, fmt.Errorf("image projection joint outside SOMA topology") + } + } + + result.joints = smpl.Mapping + return result, nil +} + +func (p poseProjection) project(camera, translation []float32) []float32 { + if len(camera) != 231 || len(translation) != 3 || p.width == 0 || p.height == 0 { + return nil + } + // gemx_live_submit documents and uses these exact source-frame intrinsics. + focal := float64(max(p.width, p.height)) + result := make([]float32, 0, len(p.joints)*2) + + for _, joint := range p.joints { + x := float64(camera[joint*3]) + float64(translation[0]) + y := float64(camera[joint*3+1]) + float64(translation[1]) + z := float64(camera[joint*3+2]) + float64(translation[2]) + if z <= 0 { + return nil + } + for _, pixel := range []float64{focal*x/z + float64(p.width)/2, focal*y/z + float64(p.height)/2} { + if math.IsNaN(pixel) || math.IsInf(pixel, 0) || math.Abs(pixel) > math.MaxFloat32 { + return nil + } + result = append(result, float32(pixel)) + } + } + return result +} diff --git a/backend/go/gemxcpp/image_projection_test.go b/backend/go/gemxcpp/image_projection_test.go new file mode 100644 index 000000000000..2362714c3e0c --- /dev/null +++ b/backend/go/gemxcpp/image_projection_test.go @@ -0,0 +1,84 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "encoding/json" + "math" + + motion "github.com/mudler/LocalAI/pkg/motion/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/protobuf/proto" +) + +var _ = Describe("Source-frame pose projection", func() { + It("projects translated camera joints using landscape and portrait frame intrinsics", func() { + camera := make([]float32, 231) + camera[3], camera[4], camera[5] = 1, -1, 2 + p := poseProjection{joints: []int{1, 0}, width: 640, height: 480} + + pixels := p.project(camera, []float32{0, 0, 2}) + + Expect(pixels).To(Equal([]float32{480, 80, 320, 240})) + + p.width, p.height = 480, 640 + pixels = p.project(camera, []float32{0, 0, 2}) + + Expect(pixels).To(Equal([]float32{400, 160, 240, 320})) + }) + It("preserves SMPL ordering from the native definition", func() { + mapping := make([]int, 24) + for i := range mapping { + mapping[i] = 76 - i + } + + raw, err := json.Marshal(map[string]any{"mapping": mapping}) + Expect(err).NotTo(HaveOccurred()) + + p, err := projectionForProfile(map[string]json.RawMessage{"smpl24": raw}, "smpl24") + + Expect(err).NotTo(HaveOccurred()) + Expect(p.joints).To(Equal(mapping)) + }) + It("preserves SOMA identity ordering", func() { + soma, err := projectionForProfile(nil, "soma77") + + Expect(err).NotTo(HaveOccurred()) + Expect(soma.joints).To(HaveLen(77)) + for i, joint := range soma.joints { + Expect(joint).To(Equal(i)) + } + }) + It("rejects an incomplete native mapping", func() { + definition := map[string]json.RawMessage{"smpl24": json.RawMessage(`{"mapping":[0]}`)} + + _, err := projectionForProfile(definition, "smpl24") + + Expect(err).To(HaveOccurred()) + }) + DescribeTable("omits image coordinates for unusable input", + func(camera []float32, translation []float32, width uint32) { + p := poseProjection{joints: []int{0}, width: width, height: 480} + + pixels := p.project(camera, translation) + + Expect(pixels).To(BeNil()) + }, + Entry("zero depth", make([]float32, 231), []float32{0, 0, 0}, uint32(640)), + Entry("behind camera", make([]float32, 231), []float32{0, 0, -1}, uint32(640)), + Entry("nonfinite translation", make([]float32, 231), []float32{float32(math.NaN()), 0, 1}, uint32(640)), + Entry("incomplete camera channel", make([]float32, 3), []float32{0, 0, 1}, uint32(640)), + Entry("missing frame dimensions", make([]float32, 231), []float32{0, 0, 1}, uint32(0)), + ) + It("round-trips optional image positions alongside frame identity", func() { + pose := &motion.Pose{Sequence: 7, SourceTimeUs: 42, ImagePositions: []float32{12, 34}} + + encoded, err := proto.Marshal(pose) + Expect(err).NotTo(HaveOccurred()) + + decoded := &motion.Pose{} + Expect(proto.Unmarshal(encoded, decoded)).To(Succeed()) + + Expect(proto.Equal(pose, decoded)).To(BeTrue()) + }) +}) diff --git a/backend/go/gemxcpp/main.go b/backend/go/gemxcpp/main.go new file mode 100644 index 000000000000..8f17fb53d25d --- /dev/null +++ b/backend/go/gemxcpp/main.go @@ -0,0 +1,22 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "flag" + "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/xlog" + "os" +) + +func main() { + addr := flag.String("addr", "localhost:50051", "gRPC listen address") + flag.Parse() + if err := loadNativeLibrary(os.Getenv("GEMX_LIBRARY")); err != nil { + xlog.Error("loading GEM-X library", "error", err) + os.Exit(1) + } + if err := grpc.StartServer(*addr, &GemX{}); err != nil { + xlog.Error("GEM-X server", "error", err) + os.Exit(1) + } +} diff --git a/backend/go/gemxcpp/native.go b/backend/go/gemxcpp/native.go new file mode 100644 index 000000000000..02a56e7c58f8 --- /dev/null +++ b/backend/go/gemxcpp/native.go @@ -0,0 +1,52 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "bytes" + "fmt" + "github.com/ebitengine/purego" +) + +var liveCreate func(string, string, string, string, string, *uint32, int64, *uintptr, *byte, uint64) int32 +var liveDestroy func(uintptr) +var liveDefinition func(uintptr, *byte, uint64, *uint64, *byte, uint64) int32 +var liveReset func(uintptr, *byte, uint64) int32 +var liveSubmit func(uintptr, *byte, uint64, uint32, uint32, uint64, uint64, int64, *float32, uint64, uint64, *uintptr, *byte, uint64) int32 +var resultDestroy func(uintptr) +var resultInfo func(uintptr, *uint64, *int64, *uint64, *uint64, *uint32, *uint32, *byte, uint64) int32 +var resultIntervalStart func(uintptr, *int64, *byte, uint64) int32 +var resultCopy func(uintptr, uint32, *float32, uint64, *uint64, *byte, uint64) int32 + +func loadNativeLibrary(path string) error { + lib, err := purego.Dlopen(path, purego.RTLD_NOW|purego.RTLD_GLOBAL) + if err != nil { + return err + } + for _, b := range []struct { + fn any + name string + }{ + {&liveCreate, "localai_gemx_create"}, {&liveDestroy, "gemx_live_destroy"}, + {&liveDefinition, "gemx_live_definition"}, {&liveReset, "gemx_live_reset"}, + {&liveSubmit, "gemx_live_submit"}, {&resultDestroy, "gemx_live_result_destroy"}, + {&resultInfo, "gemx_live_result_info"}, {&resultCopy, "gemx_live_result_copy"}, + } { + symbol, err := purego.Dlsym(lib, b.name) + if err != nil { + return fmt.Errorf("streaming ABI: %w", err) + } + purego.RegisterFunc(b.fn, symbol) + } + // Optional additive ABI; old libraries can still stream poses. + resultIntervalStart = nil + if symbol, err := purego.Dlsym(lib, "gemx_live_result_interval_start"); err == nil { + purego.RegisterFunc(&resultIntervalStart, symbol) + } + return nil +} +func nativeError(code int32, buf []byte) error { + if code == 0 { + return nil + } + return fmt.Errorf("gemx (%d): %s", code, bytes.TrimRight(buf, "\x00")) +} diff --git a/backend/go/gemxcpp/native/bridge.cpp b/backend/go/gemxcpp/native/bridge.cpp new file mode 100644 index 000000000000..bb5301a9f44e --- /dev/null +++ b/backend/go/gemxcpp/native/bridge.cpp @@ -0,0 +1,13 @@ +// SPDX-License-Identifier: MIT +#include "gemx_stream.h" + +// PureGo supports at most 15 machine arguments; bundle scalar configuration +// without asking Go callers to reproduce a platform-dependent C structure. +extern "C" GEMX_API gemx_status localai_gemx_create( + const char *gem, const char *pose, const char *detector, + const char *module, const char *backend, const uint32_t *options, + int64_t gap, gemx_live **out, char *error, uint64_t capacity) { + return gemx_live_create(gem, pose, detector, module, backend, options[0], + "", options[1], options[2], options[3], options[4], options[5], + GEMX_LIVE_FRAME_INDEX, gap, out, error, capacity); +} diff --git a/backend/go/gemxcpp/package.sh b/backend/go/gemxcpp/package.sh new file mode 100755 index 000000000000..f726842b223b --- /dev/null +++ b/backend/go/gemxcpp/package.sh @@ -0,0 +1,39 @@ +#!/bin/bash +set -euo pipefail + +BACKEND_DIR=$(cd -- "$(dirname -- "$0")" && pwd) +# The package is generated; never mix libraries from different build variants. +rm -rf -- "$BACKEND_DIR/package" +mkdir -p "$BACKEND_DIR/package/lib" +cp "$BACKEND_DIR/gemxcpp" "$BACKEND_DIR/run.sh" "$BACKEND_DIR/package/" +chmod +x "$BACKEND_DIR/package/run.sh" +find "$BACKEND_DIR/${BUILD_DIR:-build-cpu}" \( -name 'libgemx.so*' -o -name 'libgemx.dylib' -o -name 'libggml*.so*' -o -name 'libggml*.dylib' \) -exec cp -a {} "$BACKEND_DIR/package/lib/" \; +cp "$BACKEND_DIR/sources/gem-x.cpp/LICENSE" "$BACKEND_DIR/package/LICENSE.gemx" +cp "$BACKEND_DIR/sources/gem-x.cpp/NOTICE" "$BACKEND_DIR/package/NOTICE.gemx" +cp -R "$BACKEND_DIR/sources/gem-x.cpp/LICENSES" "$BACKEND_DIR/package/LICENSES.gemx" +cp "$BACKEND_DIR/sources/gem-x.cpp/ggml/LICENSE" "$BACKEND_DIR/package/LICENSE.ggml" + +if [ "$(uname -s)" != Darwin ]; then + # purego's executable also imports libdl/libpthread, even with CGO disabled. + for library in "$BACKEND_DIR/package/gemxcpp" "$BACKEND_DIR"/package/lib/*.so*; do + while read -r dependency; do + [ -f "$dependency" ] && cp -L "$dependency" "$BACKEND_DIR/package/lib/" + done < <(ldd "$library" | awk '/=> \// {print $3}') + done + loader=$(ldd "$BACKEND_DIR/package/lib/libgemx.so" | awk '/ld-linux/ {print $1; exit}') + if [ -f "$loader" ]; then cp -L "$loader" "$BACKEND_DIR/package/lib/ld.so"; fi + source "$BACKEND_DIR/../../../scripts/build/package-gpu-libs.sh" "$BACKEND_DIR/package/lib" + package_gpu_libs + if [ "${BUILD_TYPE:-}" = vulkan ]; then + # NVIDIA's host ICD dlopens EGL; it is not visible in libgemx's ldd tree. + for soname in libEGL.so.1 libGLdispatch.so.0 libX11.so.6 libXext.so.6; do + dependency=$(ldconfig -p | awk -v name="$soname" '$1 == name {print $NF; exit}') + if [ ! -f "$dependency" ]; then + echo "Missing Vulkan ICD runtime dependency: $soname" >&2 + exit 1 + fi + copy_lib "$dependency" + done + sweep_transitive_deps + fi +fi diff --git a/backend/go/gemxcpp/package_test.go b/backend/go/gemxcpp/package_test.go new file mode 100644 index 000000000000..7521d6419b5d --- /dev/null +++ b/backend/go/gemxcpp/package_test.go @@ -0,0 +1,58 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "go.yaml.in/yaml/v2" + "os" + "path/filepath" + "strings" +) + +var _ = Describe("GEM-X distribution", func() { + It("publishes every CI variant in the backend gallery", func() { + root := "../../.." + var matrix struct { + Linux []map[string]string `yaml:"include"` + Darwin []map[string]string `yaml:"includeDarwin"` + } + data, err := os.ReadFile(filepath.Join(root, ".github/backend-matrix.yml")) + Expect(err).NotTo(HaveOccurred()) + Expect(yaml.Unmarshal(data, &matrix)).To(Succeed()) + var entries []struct { + Name string `yaml:"name"` + URI string `yaml:"uri"` + Capabilities map[string]string `yaml:"capabilities"` + } + data, err = os.ReadFile(filepath.Join(root, "backend/index.yaml")) + Expect(err).NotTo(HaveOccurred()) + Expect(yaml.Unmarshal(data, &entries)).To(Succeed()) + images := map[string]bool{} + names := map[string]bool{} + var aliases map[string]string + for _, e := range entries { + images[e.URI] = true + names[e.Name] = true + if e.Name == "gemxcpp" { + aliases = e.Capabilities + } + } + count := 0 + for _, e := range append(matrix.Linux, matrix.Darwin...) { + if e["backend"] != "gemxcpp" { + continue + } + count++ + for _, version := range []string{"latest", "master"} { + Expect(images["quay.io/go-skynet/local-ai-backends:"+version+e["tag-suffix"]]).To(BeTrue()) + } + } + Expect(count).To(Equal(5)) + Expect(aliases).NotTo(BeEmpty()) + for _, name := range aliases { + Expect(names[name]).To(BeTrue()) + } + Expect(strings.HasPrefix(aliases["metal"], "cpu-darwin")).To(BeTrue()) + }) +}) diff --git a/backend/go/gemxcpp/root_displacement_test.go b/backend/go/gemxcpp/root_displacement_test.go new file mode 100644 index 000000000000..2b6abd32e89e --- /dev/null +++ b/backend/go/gemxcpp/root_displacement_test.go @@ -0,0 +1,114 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + "encoding/json" + "math" + "unsafe" + + motion "github.com/mudler/LocalAI/pkg/motion/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/protobuf/proto" +) + +var _ = Describe("Root displacement capability", func() { + BeforeEach(func() { old := resultIntervalStart; DeferCleanup(func() { resultIntervalStart = old }) }) + for _, tc := range []struct { + name, profile, channels string + symbol, want bool + }{ + {"supported", "smpl24", `{"root_displacement":3}`, true, true}, + {"old library", "smpl24", `{"root_displacement":3}`, false, false}, + {"old definition", "smpl24", `{}`, true, false}, + {"wrong shape", "smpl24", `{"root_displacement":2}`, true, false}, + {"malformed", "smpl24", `null`, true, false}, + {"SOMA basis", "soma77", `{"root_displacement":3}`, true, false}, + } { + It(tc.name, func() { + resultIntervalStart = nil + if tc.symbol { + resultIntervalStart = func(uintptr, *int64, *byte, uint64) int32 { return 0 } + } + raw := map[string]json.RawMessage{"channels": json.RawMessage(tc.channels), tc.profile: json.RawMessage(`{"schema":"test","joint_names":["root"],"parents":[-1],"root":0,"rest_local_translations":[0,0,0],"rest_local_rotations":[1,0,0,0]}`)} + d, err := publicDefinition(raw, tc.profile) + Expect(err).NotTo(HaveOccurred()) + found := false + for _, c := range d.Channels { + found = found || c == "root_displacement" + } + Expect(found).To(Equal(tc.want)) + }) + } +}) + +var _ = Describe("Root displacement results", func() { + BeforeEach(func() { + oldInfo, oldCopy, oldStart := resultInfo, resultCopy, resultIntervalStart + DeferCleanup(func() { resultInfo, resultCopy, resultIntervalStart = oldInfo, oldCopy, oldStart }) + }) + for _, tc := range []struct { + name string + start int64 + delta []float32 + enabled bool + flags, outcome uint32 + wantErr bool + }{ + {"zero timestamp", 0, []float32{1, 2, 3}, true, 0, 1, false}, + {"irregular interval", 123, []float32{-.1, 0, .2}, true, 0, 1, false}, + {"negative timestamp", -1, []float32{1, 2, 3}, true, 0, 1, true}, + {"empty interval", 1000, []float32{1, 2, 3}, true, 0, 1, true}, + {"future timestamp", 1001, []float32{1, 2, 3}, true, 0, 1, true}, + {"missing channel", 0, nil, true, 0, 1, true}, + {"short channel", 0, []float32{1, 2}, true, 0, 1, true}, + {"oversize channel", 0, []float32{1, 2, 3, 4}, true, 0, 1, true}, + {"nonfinite channel", 0, []float32{float32(math.NaN()), 0, 0}, true, 0, 1, true}, + {"legacy", 0, nil, false, 0, 1, false}, + {"reset", 0, nil, true, 1, 1, false}, + {"warmup", 0, nil, true, 1, 0, false}, + {"lost", 0, nil, true, 1, 2, false}, + } { + It(tc.name, func() { + touched := false + resultInfo = func(_ uintptr, s *uint64, ts *int64, e, tr *uint64, o, f *uint32, _ *byte, _ uint64) int32 { + *s = 2 + *ts = 1000 + *e = 3 + *tr = 4 + *o = tc.outcome + *f = tc.flags + return 0 + } + resultIntervalStart = func(_ uintptr, start *int64, _ *byte, _ uint64) int32 { touched = true; *start = tc.start; return 0 } + resultCopy = func(_ uintptr, ch uint32, out *float32, cap uint64, n *uint64, _ *byte, _ uint64) int32 { + data := map[uint32][]float32{6: make([]float32, 3), 10: make([]float32, 4), 3: make([]float32, 72), 4: {1, 0, 0, 0}}[ch] + if ch == 14 { + touched = true + data = tc.delta + } + *n = uint64(len(data)) + if out != nil { + copy(unsafe.Slice(out, int(cap)), data) + } + return 0 + } + out, _, err := readResult(1, "smpl24", poseProjection{}, tc.enabled) + Expect(err != nil).To(Equal(tc.wantErr), "%v", err) + if tc.wantErr { + return + } + want := tc.enabled && tc.flags&1 == 0 && tc.outcome == 1 + Expect(touched).To(Equal(want)) + b, err := proto.Marshal(out) + Expect(err).NotTo(HaveOccurred()) + decoded := new(motion.Output) + Expect(proto.Unmarshal(b, decoded)).To(Succeed()) + if want { + p := decoded.GetPose() + Expect(p.DisplacementStartTimeUs).To(Equal(tc.start)) + Expect(p.RootDisplacement).To(Equal(tc.delta)) + } + }) + } +}) diff --git a/backend/go/gemxcpp/run.sh b/backend/go/gemxcpp/run.sh new file mode 100755 index 000000000000..2a0f102a8ee6 --- /dev/null +++ b/backend/go/gemxcpp/run.sh @@ -0,0 +1,19 @@ +#!/bin/bash +set -euo pipefail +BACKEND_DIR=$(cd -- "$(dirname -- "$0")" && pwd) +export GEMX_MODULE="$BACKEND_DIR/lib" +export GEMX_DEVICE=CPU +if [ -f "$BACKEND_DIR/lib/libggml-vulkan.so" ]; then export GEMX_DEVICE=Vulkan; fi +# Precision must be fixed before the first process-wide GGML initialization. +export GGML_VK_DISABLE_F16=1 GGML_VK_DISABLE_COOPMAT=1 GGML_VK_DISABLE_COOPMAT2=1 +if [ "$(uname -s)" = Darwin ]; then + export GEMX_LIBRARY="$BACKEND_DIR/lib/libgemx.dylib" + export DYLD_LIBRARY_PATH="$BACKEND_DIR/lib:${DYLD_LIBRARY_PATH:-}" +else + export GEMX_LIBRARY="$BACKEND_DIR/lib/libgemx.so" + export LD_LIBRARY_PATH="$BACKEND_DIR/lib:${LD_LIBRARY_PATH:-}" +fi +if [ -f "$BACKEND_DIR/lib/ld.so" ]; then + exec "$BACKEND_DIR/lib/ld.so" "$BACKEND_DIR/gemxcpp" "$@" +fi +exec "$BACKEND_DIR/gemxcpp" "$@" diff --git a/backend/index.yaml b/backend/index.yaml index d399d6498fc0..2ca456f368ff 100644 --- a/backend/index.yaml +++ b/backend/index.yaml @@ -539,6 +539,25 @@ amd: "vulkan-kimodocpp" intel: "vulkan-kimodocpp" metal: "cpu-darwin-arm64-kimodocpp" +- &gemxcpp + name: "gemxcpp" + alias: "gemxcpp" + license: apache-2.0 + description: | + gem-x.cpp live human motion: SOMA and SMPL pose streams on CPU and Vulkan. + Apple Silicon uses the CPU runtime. + urls: + - https://github.com/localai-org/gem-x.cpp + tags: [motion, pose-estimation, CPU, Vulkan] + capabilities: + default: "cpu-gemxcpp" + vulkan: "vulkan-gemxcpp" + nvidia: "vulkan-gemxcpp" + nvidia-cuda-12: "vulkan-gemxcpp" + nvidia-cuda-13: "vulkan-gemxcpp" + amd: "vulkan-gemxcpp" + intel: "vulkan-gemxcpp" + metal: "cpu-darwin-arm64-gemxcpp" - &trellis2cpp name: "trellis2cpp" alias: "trellis2cpp" @@ -4147,6 +4166,49 @@ mirrors: - localai/localai-backends:master-cpu-darwin-arm64-kimodocpp +## gem-x.cpp +- !!merge <<: *gemxcpp + name: "gemxcpp-development" + capabilities: + default: "cpu-gemxcpp-development" + vulkan: "vulkan-gemxcpp-development" + nvidia: "vulkan-gemxcpp-development" + nvidia-cuda-12: "vulkan-gemxcpp-development" + nvidia-cuda-13: "vulkan-gemxcpp-development" + amd: "vulkan-gemxcpp-development" + intel: "vulkan-gemxcpp-development" + metal: "cpu-darwin-arm64-gemxcpp-development" +- !!merge <<: *gemxcpp + name: "cpu-gemxcpp" + uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-gemxcpp" + mirrors: + - localai/localai-backends:latest-cpu-gemxcpp +- !!merge <<: *gemxcpp + name: "cpu-gemxcpp-development" + uri: "quay.io/go-skynet/local-ai-backends:master-cpu-gemxcpp" + mirrors: + - localai/localai-backends:master-cpu-gemxcpp +- !!merge <<: *gemxcpp + name: "vulkan-gemxcpp" + uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-vulkan-gemxcpp" + mirrors: + - localai/localai-backends:latest-gpu-vulkan-gemxcpp +- !!merge <<: *gemxcpp + name: "vulkan-gemxcpp-development" + uri: "quay.io/go-skynet/local-ai-backends:master-gpu-vulkan-gemxcpp" + mirrors: + - localai/localai-backends:master-gpu-vulkan-gemxcpp +- !!merge <<: *gemxcpp + name: "cpu-darwin-arm64-gemxcpp" + uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-darwin-arm64-gemxcpp" + mirrors: + - localai/localai-backends:latest-cpu-darwin-arm64-gemxcpp +- !!merge <<: *gemxcpp + name: "cpu-darwin-arm64-gemxcpp-development" + uri: "quay.io/go-skynet/local-ai-backends:master-cpu-darwin-arm64-gemxcpp" + mirrors: + - localai/localai-backends:master-cpu-darwin-arm64-gemxcpp + ## trellis2cpp - !!merge <<: *trellis2cpp name: "cpu-trellis2cpp" diff --git a/core/application/application.go b/core/application/application.go index 326885b3105b..f67bc55eaec3 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -56,26 +56,28 @@ const faceEmbeddingDim = 0 const voiceEmbeddingDim = 0 type Application struct { - backendLoader *config.ModelConfigLoader - modelLoader *model.ModelLoader - applicationConfig *config.ApplicationConfig - startupConfig *config.ApplicationConfig // Stores original config from env vars (before file loading) - templatesEvaluator *templates.Evaluator - galleryService *galleryop.GalleryService - agentJobService *agentpool.AgentJobService - agentPoolService atomic.Pointer[agentpool.AgentPoolService] - faceRegistry facerecognition.Registry - voiceRegistry voicerecognition.Registry - voiceProfileStore *voiceprofile.Store - authDB *gorm.DB - metricsService *monitoring.LocalAIMetricsService - statsRecorder *billing.Recorder - fallbackUser *auth.User - piiRedactor *pii.Redactor - piiEvents pii.EventStore - mitmCA atomic.Pointer[mitm.CA] - mitmServer atomic.Pointer[mitm.Server] - mitmMutex sync.Mutex // serializes Stop+Start; readers use atomic loads + webSocketTicketsOnce sync.Once + webSocketTickets *auth.WebSocketTickets + backendLoader *config.ModelConfigLoader + modelLoader *model.ModelLoader + applicationConfig *config.ApplicationConfig + startupConfig *config.ApplicationConfig // Stores original config from env vars (before file loading) + templatesEvaluator *templates.Evaluator + galleryService *galleryop.GalleryService + agentJobService *agentpool.AgentJobService + agentPoolService atomic.Pointer[agentpool.AgentPoolService] + faceRegistry facerecognition.Registry + voiceRegistry voicerecognition.Registry + voiceProfileStore *voiceprofile.Store + authDB *gorm.DB + metricsService *monitoring.LocalAIMetricsService + statsRecorder *billing.Recorder + fallbackUser *auth.User + piiRedactor *pii.Redactor + piiEvents pii.EventStore + mitmCA atomic.Pointer[mitm.CA] + mitmServer atomic.Pointer[mitm.Server] + mitmMutex sync.Mutex // serializes Stop+Start; readers use atomic loads // mitmHostConflicts records duplicate-host claims across model configs. // Non-empty disables the MITM listener until resolved — the strict // 1-to-1 host↔model invariant the dispatcher relies on. Read by @@ -131,6 +133,12 @@ type Application struct { shutdownOnce sync.Once } +// WebSocketTickets returns this frontend's short-lived upgrade ticket store. +func (a *Application) WebSocketTickets() *auth.WebSocketTickets { + a.webSocketTicketsOnce.Do(func() { a.webSocketTickets = auth.NewWebSocketTickets() }) + return a.webSocketTickets +} + // Ready reports whether the application has finished starting up and can serve // traffic. It backs the /readyz probe and is safe to call from any goroutine, // including while startup is still running. diff --git a/core/backend/motion.go b/core/backend/motion.go new file mode 100644 index 000000000000..3aa49c4e033e --- /dev/null +++ b/core/backend/motion.go @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: MIT +package backend + +import ( + "context" + "fmt" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/LocalAI/pkg/model" +) + +func ModelMotionStream(ctx context.Context, loader *model.ModelLoader, app *config.ApplicationConfig, cfg config.ModelConfig) (grpc.MotionStreamClient, error) { + m, err := loader.Load(ModelOptions(cfg, app)...) + if err != nil { + recordModelLoadFailure(app, cfg.Name, cfg.Backend, err, nil) + return nil, err + } + if m == nil { + return nil, fmt.Errorf("motion model unavailable") + } + // A live session is intentionally long-lived. Its context owns the global slot. + release, err := AcquireGlobalBackendSlot() + if err != nil { + return nil, err + } + s, err := m.MotionStream(ctx) + if err != nil { + release() + return nil, err + } + go func() { <-ctx.Done(); release() }() + return s, nil +} diff --git a/core/config/backend_capabilities.go b/core/config/backend_capabilities.go index 96d53c4e255c..7bfb53ff2483 100644 --- a/core/config/backend_capabilities.go +++ b/core/config/backend_capabilities.go @@ -20,6 +20,7 @@ const ( UsecaseVideo = "video" Usecase3D = "3d" Usecase3DAnimation = "3d_animation" + UsecaseMotion = "motion" UsecaseTranscript = "transcript" UsecaseTTS = "tts" UsecaseSoundGeneration = "sound_generation" @@ -50,6 +51,7 @@ const ( MethodGenerateVideo GRPCMethod = "GenerateVideo" MethodGenerate3D GRPCMethod = "Generate3D" MethodAnimate3D GRPCMethod = "Animate3D" + MethodMotionStream GRPCMethod = "MotionStream" MethodAudioTranscription GRPCMethod = "AudioTranscription" MethodTTS GRPCMethod = "TTS" MethodTTSStream GRPCMethod = "TTSStream" @@ -90,6 +92,7 @@ type UsecaseInfo struct { // UsecaseInfoMap maps each known_usecase string to its gRPC and semantic info. var UsecaseInfoMap = map[string]UsecaseInfo{ + UsecaseMotion: {Flag: FLAG_MOTION, GRPCMethod: MethodMotionStream, Description: "Timestamped human pose streams with native or SMPL skeletons."}, UsecaseChat: { Flag: FLAG_CHAT, GRPCMethod: MethodPredict, @@ -439,6 +442,7 @@ var BackendCapabilities = map[string]BackendCapability{ }, // --- 3D generation backends --- + "gemxcpp": {GRPCMethods: []GRPCMethod{MethodMotionStream}, PossibleUsecases: []string{UsecaseMotion}, DefaultUsecases: []string{UsecaseMotion}, Description: "GEM-X live human motion on CPU/Vulkan; SOMA-77 and SMPL-24 pose streams"}, "kimodocpp": { GRPCMethods: []GRPCMethod{MethodAnimate3D}, PossibleUsecases: []string{Usecase3DAnimation}, diff --git a/core/config/backend_capabilities_test.go b/core/config/backend_capabilities_test.go index 6f683c0bf900..7c335fa15e1f 100644 --- a/core/config/backend_capabilities_test.go +++ b/core/config/backend_capabilities_test.go @@ -445,3 +445,17 @@ var _ = Describe("AllBackendNames", func() { Expect(slices.IsSorted(names)).To(BeTrue()) }) }) + +var _ = Describe("GEM-X motion capability", func() { + It("discovers motion without advertising animation generation", func() { + Expect(DefaultUsecasesForBackendCap("gemxcpp")).To(Equal([]string{UsecaseMotion})) + cfg := &ModelConfig{Backend: "gemxcpp"} + Expect(cfg.HasUsecases(FLAG_MOTION)).To(BeTrue()) + Expect(cfg.HasUsecases(FLAG_3D_ANIMATION)).To(BeFalse()) + Expect(cfg.HasUsecases(FLAG_DECISIONS)).To(BeFalse()) + decisions := FLAG_DECISIONS + Expect((&ModelConfig{Backend: "vllm-cpp", KnownUsecases: &decisions}).HasUsecases(FLAG_MOTION)).To(BeFalse()) + Expect(FLAG_MOTION & FLAG_DECISIONS).To(BeZero()) + Expect((&ModelConfig{Backend: "llama-cpp"}).HasUsecases(FLAG_MOTION)).To(BeFalse()) + }) +}) diff --git a/core/config/model_config.go b/core/config/model_config.go index b2b727628743..0667c3ceb06c 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -2080,6 +2080,7 @@ const ( // optional PBR material, e.g. trellis2cpp). FLAG_3D ModelConfigUsecase = 0b100000000000000000000000 FLAG_3D_ANIMATION ModelConfigUsecase = 1 << 24 + FLAG_MOTION ModelConfigUsecase = 1 << 26 // Marks a model as a decision model: it answers typed choice / noul / // score questions over a state (served by POST /v1/systemone). @@ -2097,7 +2098,7 @@ const ( // both text/language). A model is multimodal when its usecases span 2+ groups. var ModalityGroups = []ModelConfigUsecase{ FLAG_CHAT | FLAG_COMPLETION | FLAG_EDIT, // text/language - FLAG_VISION | FLAG_DETECTION, // visual understanding + FLAG_VISION | FLAG_DETECTION | FLAG_MOTION, // visual understanding FLAG_TRANSCRIPT | FLAG_REALTIME_AUDIO | FLAG_SOUND_CLASSIFICATION, // audio input — realtime_audio is any-to-any, so it counts here too FLAG_TTS | FLAG_SOUND_GENERATION | FLAG_REALTIME_AUDIO, // audio output — and here, so a lone realtime_audio flag still reads as multimodal FLAG_AUDIO_TRANSFORM, // audio in/out transforms @@ -2151,6 +2152,7 @@ func GetAllModelConfigUsecases() map[string]ModelConfigUsecase { "FLAG_3D": FLAG_3D, "FLAG_3D_ANIMATION": FLAG_3D_ANIMATION, "FLAG_DECISIONS": FLAG_DECISIONS, + "FLAG_MOTION": FLAG_MOTION, } } @@ -2202,6 +2204,9 @@ func (c *ModelConfig) HasUsecases(u ModelConfigUsecase) bool { // In its current state, this function should ideally check for properties of the config like templates, rather than the direct backend name checks for the lower half. // This avoids the maintenance burden of updating this list for each new backend - but unfortunately, that's the best option for some services currently. func (c *ModelConfig) GuessUsecases(u ModelConfigUsecase) bool { + if u&FLAG_MOTION != 0 && c.Backend != "gemxcpp" { + return false + } // Backends that are clearly not text-generation nonTextGenBackends := []string{ "whisper", "piper", "kokoro", diff --git a/core/gallery/importers/gemxcpp.go b/core/gallery/importers/gemxcpp.go new file mode 100644 index 000000000000..5276737f5ff7 --- /dev/null +++ b/core/gallery/importers/gemxcpp.go @@ -0,0 +1,70 @@ +// SPDX-License-Identifier: MIT +package importers + +import ( + "encoding/json" + "fmt" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/gallery" + "github.com/mudler/LocalAI/core/schema" + "go.yaml.in/yaml/v2" + "strings" +) + +type GemXCppImporter struct{} + +func (*GemXCppImporter) Name() string { return "gemxcpp" } +func (*GemXCppImporter) Modality() string { return "motion" } +func (*GemXCppImporter) AutoDetects() bool { return true } +func gemxRepository(uri string) bool { + owner, repo, ok := HFOwnerRepoFromURI(uri) + return ok && strings.EqualFold(owner+"/"+repo, "LocalAI-io/GEM-X-GGUF") +} +func (*GemXCppImporter) Match(d Details) bool { + var p struct { + Backend string `json:"backend"` + } + if len(d.Preferences) > 0 { + if json.Unmarshal(d.Preferences, &p) != nil { + return false + } + } + if p.Backend != "" { + return p.Backend == "gemxcpp" + } + return gemxRepository(d.URI) +} +func (*GemXCppImporter) Import(d Details) (gallery.ModelConfig, error) { + var p struct { + Name string `json:"name"` + Description string `json:"description"` + } + if len(d.Preferences) > 0 { + if err := json.Unmarshal(d.Preferences, &p); err != nil { + return gallery.ModelConfig{}, err + } + } + if !gemxRepository(d.URI) { + return gallery.ModelConfig{}, fmt.Errorf("gemxcpp requires the published LocalAI-io/GEM-X-GGUF bundle") + } + if p.Name == "" { + p.Name = "gem-x" + } + if p.Description == "" { + p.Description = "Live human motion capture: SOMA-77 or SMPL-24 poses on CPU/Vulkan" + } + files := []gallery.File{} + for _, f := range []struct{ name, sha string }{ + {"gem-x-contact-f32.gguf", "175857b8453b65d028f3703d813c3186b31b66ea65fbdb94e64308c39e083f78"}, + {"vitpose-f32.gguf", "272c75d4c3a6a740f1eb3f3222832de206150fa035c134af17e66ca9f8386d11"}, + {"yolox-f32.gguf", "2be3d28e0dd8a171f4ad980b3dfbddaa1ecb88e478c612bc11379913908e7b78"}, + } { + files = append(files, gallery.File{Filename: "gem-x/" + f.name, URI: "https://huggingface.co/LocalAI-io/GEM-X-GGUF/resolve/b36180fd4c0d7c6a7fb2269348f632d3ed279868/" + f.name, SHA256: f.sha}) + } + cfg := config.ModelConfig{Name: p.Name, Description: p.Description, Backend: "gemxcpp", KnownUsecaseStrings: []string{"motion"}, Options: []string{"vitpose:gem-x/vitpose-f32.gguf", "yolox:gem-x/yolox-f32.gguf", "selection:continuity", "window:30", "detector_interval:1"}, PredictionOptions: schema.PredictionOptions{BasicModelRequest: schema.BasicModelRequest{Model: "gem-x/gem-x-contact-f32.gguf"}}} + data, err := yaml.Marshal(cfg) + if err != nil { + return gallery.ModelConfig{}, err + } + return gallery.ModelConfig{Name: p.Name, Description: p.Description, Files: files, ConfigFile: string(data)}, nil +} diff --git a/core/gallery/importers/gemxcpp_test.go b/core/gallery/importers/gemxcpp_test.go new file mode 100644 index 000000000000..f8a2f5080e3d --- /dev/null +++ b/core/gallery/importers/gemxcpp_test.go @@ -0,0 +1,45 @@ +// SPDX-License-Identifier: MIT +package importers_test + +import ( + "github.com/mudler/LocalAI/core/gallery" + "github.com/mudler/LocalAI/core/gallery/importers" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "go.yaml.in/yaml/v2" + "os" +) + +var _ = Describe("GemXCppImporter", func() { + It("detects the published bundle without capturing unrelated GGUFs", func() { + i := &importers.GemXCppImporter{} + Expect(i.Match(importers.Details{URI: "https://huggingface.co/LocalAI-io/GEM-X-GGUF"})).To(BeTrue()) + Expect(i.Match(importers.Details{URI: "https://huggingface.co/other/llama-GGUF"})).To(BeFalse()) + Expect(i.Match(importers.Details{URI: "https://huggingface.co/LocalAI-io/GEM-X-GGUF", Preferences: []byte(`{"backend":"llama-cpp"}`)})).To(BeFalse()) + Expect(i.Match(importers.Details{URI: "unknown", Preferences: []byte(`{"backend":"gemxcpp"}`)})).To(BeTrue()) + _, err := i.Import(importers.Details{URI: "unknown"}) + Expect(err).To(HaveOccurred()) + }) + It("keeps gallery and importer pinned component inventories identical", func() { + model, err := (&importers.GemXCppImporter{}).Import(importers.Details{URI: "https://huggingface.co/LocalAI-io/GEM-X-GGUF", Preferences: []byte(`{"name":"custom-motion"}`)}) + Expect(err).NotTo(HaveOccurred()) + Expect(model.Name).To(Equal("custom-motion")) + Expect(model.ConfigFile).To(ContainSubstring("gemxcpp")) + Expect(model.Files).To(HaveLen(3)) + data, err := os.ReadFile("../../../gallery/index.yaml") + Expect(err).NotTo(HaveOccurred()) + var entries []struct { + Name string `yaml:"name"` + Files []gallery.File `yaml:"files"` + } + Expect(yaml.Unmarshal(data, &entries)).To(Succeed()) + found := false + for _, e := range entries { + if e.Name == "gem-x" { + found = true + Expect(e.Files).To(Equal(model.Files)) + } + } + Expect(found).To(BeTrue()) + }) +}) diff --git a/core/gallery/importers/importers.go b/core/gallery/importers/importers.go index 59ecf4fccb77..ce1a14055c78 100644 --- a/core/gallery/importers/importers.go +++ b/core/gallery/importers/importers.go @@ -149,6 +149,7 @@ var defaultImporters = []Importer{ // generic .gguf importer; matches only trellis-named URIs/repos or the // distinctive component filenames, so arbitrary GGUFs are never claimed. &Trellis2CppImporter{}, + &GemXCppImporter{}, &KimodoCppImporter{}, &ACEStepImporter{}, // LongCat repositories carry generic Diffusers metadata, so this exact diff --git a/core/http/app.go b/core/http/app.go index b6211686baf4..cd63f2e6ccb0 100644 --- a/core/http/app.go +++ b/core/http/app.go @@ -374,7 +374,8 @@ func API(application *application.Application) (*echo.Echo, error) { // Build auth middleware: use the new auth.Middleware when auth is enabled or // as a unified replacement for the legacy key-auth middleware. - authMiddleware := auth.Middleware(application.AuthDB(), application.ApplicationConfig()) + authMiddleware := auth.WithWebSocketTickets(application.WebSocketTickets(), + auth.Middleware(application.AuthDB(), application.ApplicationConfig())) // Favicon handler e.GET("/favicon.svg", func(c echo.Context) error { @@ -457,7 +458,8 @@ func API(application *application.Application) (*echo.Echo, error) { // could never read a token to send back. if !application.ApplicationConfig().DisableCSRF { xlog.Debug("Enabling CSRF middleware (Sec-Fetch-Site mode)") - e.Use(auth.CSRFMiddleware()) + e.Use(auth.CSRFMiddlewareWithCORS( + application.ApplicationConfig().CORS && application.ApplicationConfig().CORSAllowOrigins != "")) } // Admin middleware: enforces admin role when auth is enabled, no-op otherwise diff --git a/core/http/auth/csrf.go b/core/http/auth/csrf.go index d768642c1c5f..c288436ecd5f 100644 --- a/core/http/auth/csrf.go +++ b/core/http/auth/csrf.go @@ -10,6 +10,12 @@ import ( // CSRFMiddleware must run after Middleware so only validated header credentials // grant an exemption. Cookie authentication must still pass the browser checks. func CSRFMiddleware() echo.MiddlewareFunc { + return CSRFMiddlewareWithCORS(false) +} + +// CSRFMiddlewareWithCORS also trusts origins approved by an explicitly configured CORS policy. +// CORS middleware must run before this middleware. +func CSRFMiddlewareWithCORS(explicitCORS bool) echo.MiddlewareFunc { return middleware.CSRFWithConfig(middleware.CSRFConfig{ Skipper: func(c echo.Context) bool { if authenticated, _ := c.Get(contextKeyHeaderAuthenticated).(bool); authenticated { @@ -19,7 +25,12 @@ func CSRFMiddleware() echo.MiddlewareFunc { return c.Request().Header.Get("Sec-Fetch-Site") == "" }, AllowSecFetchSiteFunc: func(c echo.Context) (bool, error) { - return c.Request().Header.Get("Sec-Fetch-Site") == "same-site", nil + if c.Request().Header.Get("Sec-Fetch-Site") == "same-site" { + return true, nil + } + origin := c.Request().Header.Get(echo.HeaderOrigin) + allowed := c.Response().Header().Get(echo.HeaderAccessControlAllowOrigin) + return explicitCORS && origin != "" && (allowed == "*" || allowed == origin), nil }, }) } diff --git a/core/http/auth/features.go b/core/http/auth/features.go index 4c1f53ec271e..6d93c64c8ffa 100644 --- a/core/http/auth/features.go +++ b/core/http/auth/features.go @@ -10,6 +10,11 @@ type RouteFeature struct { // RouteFeatureRegistry is the single source of truth for endpoint -> feature mappings. // To gate a new endpoint, add an entry here -- no other file changes needed. var RouteFeatureRegistry = []RouteFeature{ + {"POST", "/api/motion/sessions", FeatureMotion}, + {"GET", "/api/motion/sessions/:id", FeatureMotion}, + {"DELETE", "/api/motion/sessions/:id", FeatureMotion}, + {"GET", "/api/motion/sessions/:id/poses", FeatureMotion}, + {"POST", "/api/motion/sessions/:id/tickets", FeatureMotion}, // Chat / Completions {"POST", "/v1/chat/completions", FeatureChat}, {"POST", "/chat/completions", FeatureChat}, @@ -202,6 +207,7 @@ func APIFeatureMetas() []FeatureMeta { {FeatureDetection, "Detection", true}, {FeatureVideo, "Video Generation", true}, {Feature3D, "3D Generation", true}, + {FeatureMotion, "Motion Capture", true}, {FeatureEmbeddings, "Embeddings", true}, {FeatureSound, "Sound Generation", true}, {FeatureRealtime, "Realtime", true}, diff --git a/core/http/auth/features_motion_test.go b/core/http/auth/features_motion_test.go new file mode 100644 index 000000000000..940abaf2d2d4 --- /dev/null +++ b/core/http/auth/features_motion_test.go @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: MIT +package auth_test + +import ( + . "github.com/mudler/LocalAI/core/http/auth" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Motion feature registration", func() { + It("gates every session, ingestion and subscription route", func() { + routes := []string{} + for _, r := range RouteFeatureRegistry { + if r.Feature == FeatureMotion { + routes = append(routes, r.Method+" "+r.Pattern) + } + } + Expect(routes).To(ConsistOf("POST /api/motion/sessions", "GET /api/motion/sessions/:id", "DELETE /api/motion/sessions/:id", "GET /api/motion/sessions/:id/poses", "POST /api/motion/sessions/:id/tickets")) + Expect(APIFeatures).To(ContainElement(FeatureMotion)) + }) +}) diff --git a/core/http/auth/header_auth_test.go b/core/http/auth/header_auth_test.go new file mode 100644 index 000000000000..d5f8ef94be77 --- /dev/null +++ b/core/http/auth/header_auth_test.go @@ -0,0 +1,62 @@ +//go:build auth + +package auth_test + +import ( + "net/http" + "net/http/httptest" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Validated header authentication", func() { + DescribeTable("distinguishes supplied credentials from cookie authentication", func(mode string, expected bool) { + db := testDB() + user := createTestUser(db, "header@example.test", auth.RoleUser, auth.ProviderLocal) + session := createTestSession(db, user.ID) + key, _, err := auth.CreateAPIKey(db, user.ID, "header-test", auth.RoleUser, "", nil) + Expect(err).NotTo(HaveOccurred()) + cfg := config.NewApplicationConfig() + cfg.Auth.Enabled = true + e := echo.New() + e.Use(auth.Middleware(db, cfg)) + e.GET("/api/motion/check", func(c echo.Context) error { + Expect(auth.HeaderAuthenticated(c)).To(Equal(expected)) + return c.NoContent(204) + }) + r := httptest.NewRequest("GET", "/api/motion/check", nil) + switch mode { + case "session bearer": + r.Header.Set("Authorization", "Bearer "+session) + case "key bearer": + r.Header.Set("Authorization", "Bearer "+key) + case "key header": + r.Header.Set("x-api-key", key) + case "session cookie": + r.AddCookie(&http.Cookie{Name: "session", Value: session}) + case "key cookie": + r.AddCookie(&http.Cookie{Name: "token", Value: key}) + case "cookie with forged header": + r.AddCookie(&http.Cookie{Name: "session", Value: session}) + r.Header.Set("Authorization", "Bearer invalid") + case "key cookie with forged header": + r.AddCookie(&http.Cookie{Name: "token", Value: key}) + r.Header.Set("Authorization", "Bearer invalid") + } + res := httptest.NewRecorder() + e.ServeHTTP(res, r) + Expect(res.Code).To(Equal(204)) + }, + Entry("session bearer", "session bearer", true), + Entry("API key bearer", "key bearer", true), + Entry("API key header", "key header", true), + Entry("session cookie", "session cookie", false), + Entry("API key cookie", "key cookie", false), + Entry("cookie with forged header", "cookie with forged header", false), + Entry("key cookie with forged header", "key cookie with forged header", false), + ) +}) diff --git a/core/http/auth/middleware.go b/core/http/auth/middleware.go index f7920d526212..8889bfe57616 100644 --- a/core/http/auth/middleware.go +++ b/core/http/auth/middleware.go @@ -79,6 +79,17 @@ func Middleware(db *gorm.DB, appConfig *config.ApplicationConfig) echo.Middlewar c.Set(contextKeyUser, syntheticUser) c.Set(contextKeyRole, RoleAdmin) c.Set(contextKeySource, UsageSourceLegacy) + for _, name := range []string{"Authorization", "X-Api-Key", "Xi-Api-Key"} { + if value := c.Request().Header.Get(name); value != "" { + c.Set(contextKeyTicketCredential, ticketCredential{name, value}) + break + } + } + if c.Get(contextKeyTicketCredential) == nil { + if cookie, err := c.Cookie("token"); err == nil { + c.Set(contextKeyTicketCredential, ticketCredential{"Cookie", cookie.String()}) + } + } c.Set(contextKeyHeaderAuthenticated, extractHeaderKey(c) != "") authenticated = true } @@ -459,6 +470,7 @@ func tryAuthenticate(c echo.Context, db *gorm.DB, appConfig *config.ApplicationC if user, session := ValidateSession(db, cookie.Value, hmacSecret); user != nil { // Store session for rotation check in middleware c.Set("_auth_session", session) + c.Set(contextKeyTicketCredential, ticketCredential{"Cookie", cookie.String()}) c.Set(contextKeySource, UsageSourceWeb) return user } @@ -472,6 +484,7 @@ func tryAuthenticate(c echo.Context, db *gorm.DB, appConfig *config.ApplicationC // b1. Session token via Bearer -> still web UI if user, _ := ValidateSession(db, token, hmacSecret); user != nil { c.Set(contextKeyHeaderAuthenticated, true) + c.Set(contextKeyTicketCredential, ticketCredential{"Authorization", authHeader}) c.Set(contextKeySource, UsageSourceWeb) return user } @@ -479,6 +492,7 @@ func tryAuthenticate(c echo.Context, db *gorm.DB, appConfig *config.ApplicationC // b2. Named API key if key, err := ValidateAPIKey(db, token, hmacSecret); err == nil { c.Set(contextKeyHeaderAuthenticated, true) + c.Set(contextKeyTicketCredential, ticketCredential{"Authorization", authHeader}) c.Set(contextKeySource, UsageSourceAPIKey) c.Set(contextKeyAPIKey, key) return &key.User @@ -490,6 +504,7 @@ func tryAuthenticate(c echo.Context, db *gorm.DB, appConfig *config.ApplicationC if k := c.Request().Header.Get(header); k != "" { if apiKey, err := ValidateAPIKey(db, k, hmacSecret); err == nil { c.Set(contextKeyHeaderAuthenticated, true) + c.Set(contextKeyTicketCredential, ticketCredential{header, k}) c.Set(contextKeySource, UsageSourceAPIKey) c.Set(contextKeyAPIKey, apiKey) return &apiKey.User @@ -501,6 +516,7 @@ func tryAuthenticate(c echo.Context, db *gorm.DB, appConfig *config.ApplicationC if cookie, err := c.Cookie("token"); err == nil && cookie.Value != "" { if key, err := ValidateAPIKey(db, cookie.Value, hmacSecret); err == nil { c.Set(contextKeySource, UsageSourceAPIKey) + c.Set(contextKeyTicketCredential, ticketCredential{"Cookie", cookie.String()}) c.Set(contextKeyAPIKey, key) return &key.User } @@ -583,3 +599,9 @@ func authError(c echo.Context, appConfig *config.ApplicationConfig) error { }, }) } + +// HeaderAuthenticated reports validation of an explicitly supplied credential. +func HeaderAuthenticated(c echo.Context) bool { + validated, _ := c.Get(contextKeyHeaderAuthenticated).(bool) + return validated && GetUser(c) != nil +} diff --git a/core/http/auth/motion_permissions_test.go b/core/http/auth/motion_permissions_test.go new file mode 100644 index 000000000000..d27314443ba1 --- /dev/null +++ b/core/http/auth/motion_permissions_test.go @@ -0,0 +1,86 @@ +//go:build auth + +// SPDX-License-Identifier: MIT +package auth_test + +import ( + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "net/http" + "net/http/httptest" +) + +var _ = Describe("Motion feature authorization", func() { + It("denies revoked feature access before a WebSocket upgrade", func() { + db := testDB() + user := createTestUser(db, "motion@example.test", auth.RoleUser, auth.ProviderLocal) + token := createTestSession(db, user.ID) + Expect(auth.UpdateUserPermissions(db, user.ID, auth.PermissionMap{auth.FeatureMotion: false})).To(Succeed()) + e := newAuthTestApp(db, &config.ApplicationConfig{}) + e.GET("/api/motion/sessions/:id/poses", func(c echo.Context) error { return c.NoContent(204) }) + request := func() int { + r := httptest.NewRequest(http.MethodGet, "/api/motion/sessions/test/poses", nil) + r.Header.Set("Authorization", "Bearer "+token) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, r) + return rec.Code + } + Expect(request()).To(Equal(403)) + Expect(auth.UpdateUserPermissions(db, user.ID, auth.PermissionMap{auth.FeatureMotion: true})).To(Succeed()) + Expect(request()).To(Equal(204)) + }) +}) + +var _ = Describe("WebSocket ticket authorization", func() { + It("rechecks feature permission and session revocation at redemption", func() { + db := testDB() + user := createTestUser(db, "ticket@example.test", auth.RoleUser, auth.ProviderLocal) + token := createTestSession(db, user.ID) + alternateToken := createTestSession(db, user.ID) + cfg := config.NewApplicationConfig() + tickets := auth.NewWebSocketTickets() + e := echo.New() + e.Use(auth.WithWebSocketTickets(tickets, auth.Middleware(db, cfg))) + e.Use(auth.RequireRouteFeature(db)) + var issued auth.WebSocketTicketResponse + e.POST("/api/motion/sessions/:id/tickets", func(c echo.Context) error { + var err error + issued, err = tickets.Issue(c, "/api/motion/sessions/test/poses", "https://consumer.example") + if err != nil { + return err + } + return c.NoContent(201) + }) + e.GET("/api/motion/sessions/:id/poses", func(c echo.Context) error { return c.NoContent(204) }) + issue := func() { + r := httptest.NewRequest("POST", "/api/motion/sessions/test/tickets", nil) + // The cookie authenticates first. The unrelated valid bearer must not + // be retained as a fallback when that cookie session is revoked. + r.AddCookie(&http.Cookie{Name: "session", Value: token}) + r.Header.Set("Authorization", "Bearer "+alternateToken) + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + Expect(w.Code).To(Equal(201)) + } + redeem := func() int { + r := httptest.NewRequest("GET", "/api/motion/sessions/test/poses", nil) + r.Header.Set("Origin", "https://consumer.example") + r.Header.Set("Connection", "Upgrade") + r.Header.Set("Upgrade", "websocket") + r.Header.Set("Sec-WebSocket-Protocol", "localai.motion.v2, "+auth.WebSocketTicketProtocolPrefix+issued.Ticket) + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + return w.Code + } + issue() + Expect(auth.UpdateUserPermissions(db, user.ID, auth.PermissionMap{auth.FeatureMotion: false})).To(Succeed()) + Expect(redeem()).To(Equal(403)) + Expect(auth.UpdateUserPermissions(db, user.ID, auth.PermissionMap{auth.FeatureMotion: true})).To(Succeed()) + issue() + Expect(auth.DeleteSession(db, token, "")).To(Succeed()) + Expect(redeem()).To(Equal(401)) + }) +}) diff --git a/core/http/auth/permissions.go b/core/http/auth/permissions.go index 3f8b6deff7ed..58ccaaa85679 100644 --- a/core/http/auth/permissions.go +++ b/core/http/auth/permissions.go @@ -48,6 +48,7 @@ const ( FeatureDetection = "detection" FeatureVideo = "video" Feature3D = "3d" + FeatureMotion = "motion" FeatureEmbeddings = "embeddings" FeatureSound = "sound" FeatureRealtime = "realtime" @@ -76,7 +77,7 @@ var GeneralFeatures = []string{FeatureFineTuning, FeatureQuantization} var APIFeatures = []string{ FeatureChat, FeatureImages, FeatureAudioSpeech, FeatureAudioTranscription, FeatureAudioDiarization, FeatureAudioClassification, - FeatureVAD, FeatureDetection, FeatureVideo, Feature3D, FeatureEmbeddings, FeatureSound, + FeatureVAD, FeatureDetection, FeatureVideo, Feature3D, FeatureMotion, FeatureEmbeddings, FeatureSound, FeatureRealtime, FeatureModeration, FeatureRerank, FeatureTokenize, FeatureMCP, FeatureStores, FeatureFaceRecognition, FeatureVoiceRecognition, FeatureAudioTransform, FeaturePIIFilter, FeatureDecisions, diff --git a/core/http/auth/websocket_tickets.go b/core/http/auth/websocket_tickets.go new file mode 100644 index 000000000000..6176ac668dca --- /dev/null +++ b/core/http/auth/websocket_tickets.go @@ -0,0 +1,200 @@ +package auth + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "net/http" + "net/url" + "strings" + "sync" + "time" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" +) + +const WebSocketTicketProtocolPrefix = "localai.ticket." +const WebSocketTicketTTL = 30 * time.Second +const contextKeyWebSocketTicket = "auth_websocket_ticket" +const contextKeyTicketCredential = "auth_ticket_credential" + +type ticketCredential struct{ name, value string } + +// WebSocketTickets is application-local: tickets cannot move between frontends. +// Only hashes of random tickets are retained. Credentials are retained briefly +// so the normal authentication path can recheck revocation at redemption. +type WebSocketTickets struct { + mu sync.Mutex + entries map[[32]byte]webSocketTicket + now func() time.Time +} +type webSocketTicket struct { + path, origin, owner string + credentials http.Header + expires time.Time + timer *time.Timer +} +type WebSocketTicketResponse struct { + Ticket string `json:"ticket"` + ExpiresAt time.Time `json:"expires_at"` +} + +func NewWebSocketTickets() *WebSocketTickets { + return &WebSocketTickets{entries: make(map[[32]byte]webSocketTicket), now: time.Now} +} +func ticketOwner(c echo.Context) string { + if user := GetUser(c); user != nil { + return user.ID + } + return "" +} +func ValidWebSocketOrigin(origin string) bool { + u, err := url.Parse(origin) + return err == nil && (u.Scheme == "https" || u.Scheme == "http") && u.Host != "" && u.User == nil && u.Path == "" && u.RawQuery == "" && !u.ForceQuery && u.Fragment == "" && !strings.ContainsAny(origin, "\r\n ,") +} +func (s *WebSocketTickets) expire(now time.Time) { + for hash, ticket := range s.entries { + if !now.Before(ticket.expires) { + if ticket.timer != nil { + ticket.timer.Stop() + } + delete(s.entries, hash) + } + } +} + +// Issue must be called only after the endpoint authorizes the target resource. +func (s *WebSocketTickets) Issue(c echo.Context, path, origin string) (WebSocketTicketResponse, error) { + if !ValidWebSocketOrigin(origin) { + return WebSocketTicketResponse{}, echo.NewHTTPError(400, "a valid HTTP(S) browser origin is required") + } + if requestOrigin := c.Request().Header.Get("Origin"); requestOrigin != "" && requestOrigin != origin { + return WebSocketTicketResponse{}, echo.NewHTTPError(400, "ticket origin must match the requesting browser origin") + } + credentials := make(http.Header) + if credential, ok := c.Get(contextKeyTicketCredential).(ticketCredential); ok { + credentials.Set(credential.name, credential.value) + } + s.mu.Lock() + defer s.mu.Unlock() + now := s.now() + s.expire(now) + owner := ticketOwner(c) + count := 0 + for _, ticket := range s.entries { + if ticket.owner == owner { + count++ + } + } + if len(s.entries) >= 4096 || count >= 32 { + return WebSocketTicketResponse{}, echo.NewHTTPError(429, "too many outstanding WebSocket tickets") + } + raw := make([]byte, 32) + if _, err := rand.Read(raw); err != nil { + return WebSocketTicketResponse{}, err + } + token := base64.RawURLEncoding.EncodeToString(raw) + expires := now.Add(WebSocketTicketTTL) + hash := sha256.Sum256([]byte(token)) + ticket := webSocketTicket{path: path, origin: origin, owner: owner, credentials: credentials, expires: expires} + // Bound credential lifetime even when no subsequent requests arrive. + ticket.timer = time.AfterFunc(WebSocketTicketTTL, func() { s.mu.Lock(); delete(s.entries, hash); s.mu.Unlock() }) + s.entries[hash] = ticket + return WebSocketTicketResponse{Ticket: token, ExpiresAt: expires}, nil +} +func (s *WebSocketTickets) consume(token, path, origin string) (webSocketTicket, bool) { + s.mu.Lock() + defer s.mu.Unlock() + s.expire(s.now()) + hash := sha256.Sum256([]byte(token)) + ticket, ok := s.entries[hash] + if !ok || ticket.path != path || ticket.origin != origin { + return webSocketTicket{}, false + } + delete(s.entries, hash) + if ticket.timer != nil { + ticket.timer.Stop() + } + return ticket, true +} + +// RevokePath releases outstanding tickets when the target resource closes. +func (s *WebSocketTickets) RevokePath(path string) { + s.mu.Lock() + defer s.mu.Unlock() + for hash, ticket := range s.entries { + if ticket.path == path { + if ticket.timer != nil { + ticket.timer.Stop() + } + delete(s.entries, hash) + } + } +} + +func WebSocketTicketAuthenticated(c echo.Context) bool { + ok, _ := c.Get(contextKeyWebSocketTicket).(bool) + return ok +} + +// WithWebSocketTickets authenticates a scoped ticket through the normal auth +// middleware, retaining feature, model and ownership checks on every upgrade. +func WithWebSocketTickets(store *WebSocketTickets, authenticate echo.MiddlewareFunc) echo.MiddlewareFunc { + return func(next echo.HandlerFunc) echo.HandlerFunc { + ordinary := authenticate(next) + return func(c echo.Context) error { + var token string + protocols := []string{} + ticketCount := 0 + for _, protocol := range websocket.Subprotocols(c.Request()) { + if strings.HasPrefix(protocol, WebSocketTicketProtocolPrefix) { + ticketCount++ + token = strings.TrimPrefix(protocol, WebSocketTicketProtocolPrefix) + } else { + protocols = append(protocols, protocol) + } + } + if ticketCount == 0 { + return ordinary(c) + } + // Remove the credential before any subsequent request logging or upgrade. + c.Request().Header.Del("Sec-WebSocket-Protocol") + if len(protocols) > 0 { + c.Request().Header.Set("Sec-WebSocket-Protocol", strings.Join(protocols, ", ")) + } + if ticketCount != 1 || len(token) != 43 || c.Request().Method != http.MethodGet || !websocket.IsWebSocketUpgrade(c.Request()) { + return echo.NewHTTPError(401, "invalid WebSocket ticket") + } + ticket, ok := store.consume(token, c.Request().URL.Path, c.Request().Header.Get("Origin")) + if !ok { + return echo.NewHTTPError(401, "invalid or expired WebSocket ticket") + } + original := c.Request() + request := original.Clone(original.Context()) + for _, name := range []string{"Authorization", "X-Api-Key", "Xi-Api-Key", "Cookie"} { + request.Header.Del(name) + for _, value := range ticket.credentials.Values(name) { + request.Header.Add(name, value) + } + } + owner := ticket.owner + ticket.credentials = nil + c.SetRequest(request) + defer c.SetRequest(original) + return authenticate(func(c echo.Context) error { + if ticketOwner(c) != owner { + return echo.NewHTTPError(401, "WebSocket ticket authentication is no longer valid") + } + // Do not retain the original credential for the WebSocket lifetime. + for _, name := range []string{"Authorization", "X-Api-Key", "Xi-Api-Key", "Cookie"} { + request.Header.Del(name) + } + c.Set(contextKeyTicketCredential, nil) + c.SetRequest(original) + c.Set(contextKeyWebSocketTicket, true) + return next(c) + })(c) + } + } +} diff --git a/core/http/auth/websocket_tickets_internal_test.go b/core/http/auth/websocket_tickets_internal_test.go new file mode 100644 index 000000000000..b54da1aacfef --- /dev/null +++ b/core/http/auth/websocket_tickets_internal_test.go @@ -0,0 +1,88 @@ +package auth + +import ( + "net/http/httptest" + "sync" + "sync/atomic" + "time" + + "github.com/labstack/echo/v4" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("WebSocket ticket store", func() { + var store *WebSocketTickets + var c echo.Context + const path = "/api/motion/sessions/test/poses" + const origin = "https://consumer.example" + BeforeEach(func() { + store = NewWebSocketTickets() + c = echo.New().NewContext(httptest.NewRequest("POST", "/tickets", nil), httptest.NewRecorder()) + }) + It("expires after 30 seconds", func() { + now := time.Now() + store.now = func() time.Time { return now } + ticket, err := store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + Expect(ticket.ExpiresAt).To(Equal(now.Add(30 * time.Second))) + now = now.Add(30 * time.Second) + _, ok := store.consume(ticket.Ticket, path, origin) + Expect(ok).To(BeFalse()) + }) + It("binds the exact origin and path and consumes once", func() { + ticket, err := store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + _, ok := store.consume(ticket.Ticket, path, "https://other.example") + Expect(ok).To(BeFalse()) + _, ok = store.consume(ticket.Ticket, "/api/motion/sessions/other/poses", origin) + Expect(ok).To(BeFalse()) + _, ok = store.consume(ticket.Ticket, path, origin) + Expect(ok).To(BeTrue()) + _, ok = store.consume(ticket.Ticket, path, origin) + Expect(ok).To(BeFalse()) + }) + It("has one winner under concurrent redemption", func() { + ticket, err := store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + var wg sync.WaitGroup + var accepted atomic.Int64 + for range 16 { + wg.Go(func() { + if _, ok := store.consume(ticket.Ticket, path, origin); ok { + accepted.Add(1) + } + }) + } + wg.Wait() + Expect(accepted.Load()).To(Equal(int64(1))) + }) + It("limits outstanding tickets and frees expired capacity", func() { + now := time.Now() + store.now = func() time.Time { return now } + for range 32 { + _, err := store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + } + _, err := store.Issue(c, path, origin) + Expect(err).To(MatchError(echo.NewHTTPError(429, "too many outstanding WebSocket tickets"))) + now = now.Add(WebSocketTicketTTL) + _, err = store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + }) + It("revokes tickets when a target closes", func() { + ticket, err := store.Issue(c, path, origin) + Expect(err).NotTo(HaveOccurred()) + store.RevokePath(path) + _, ok := store.consume(ticket.Ticket, path, origin) + Expect(ok).To(BeFalse()) + }) + It("rejects mismatched request origins and opaque origins", func() { + c.Request().Header.Set("Origin", origin) + _, err := store.Issue(c, path, "https://other.example") + Expect(err).To(HaveOccurred()) + for _, invalid := range []string{"null", "*", "https://user:pass@example.com", "https://example.com/path", "https://example.com?x=y"} { + Expect(ValidWebSocketOrigin(invalid)).To(BeFalse()) + } + }) +}) diff --git a/core/http/endpoints/localai/api_instructions.go b/core/http/endpoints/localai/api_instructions.go index 06295b0ef374..0c8f4619be21 100644 --- a/core/http/endpoints/localai/api_instructions.go +++ b/core/http/endpoints/localai/api_instructions.go @@ -24,6 +24,7 @@ type instructionDef struct { } var instructionDefs = []instructionDef{ + {Name: "motion", Description: "Live human pose capture with Protobuf over WebSocket.", Tags: []string{"motion"}, Intro: "Create /api/motion/sessions and negotiate localai.motion.v2 on /poses for binary Protobuf RGB uploads and pose outputs with JSON flow credits. Browsers can POST /api/motion/sessions/{id}/tickets for a single-use 30-second WebSocket ticket. No SONIC controller or offline export."}, { Name: "chat-inference", Description: "OpenAI-compatible chat completions, text completions, and embeddings", diff --git a/core/http/endpoints/localai/api_instructions_test.go b/core/http/endpoints/localai/api_instructions_test.go index a3b504304b4f..54ed7c21e22c 100644 --- a/core/http/endpoints/localai/api_instructions_test.go +++ b/core/http/endpoints/localai/api_instructions_test.go @@ -39,7 +39,7 @@ var _ = Describe("API Instructions Endpoints", func() { instructions, ok := resp["instructions"].([]any) Expect(ok).To(BeTrue()) - Expect(instructions).To(HaveLen(21)) + Expect(instructions).To(HaveLen(22)) // Verify each instruction has required fields and correct URL format for _, s := range instructions { @@ -83,6 +83,7 @@ var _ = Describe("API Instructions Endpoints", func() { "3d", "failover", "decisions", + "motion", )) }) }) diff --git a/core/http/endpoints/localai/motion.go b/core/http/endpoints/localai/motion.go new file mode 100644 index 000000000000..bb68b01b53ce --- /dev/null +++ b/core/http/endpoints/localai/motion.go @@ -0,0 +1,414 @@ +// SPDX-License-Identifier: MIT +package localai + +import ( + "context" + "math" + "net/http" + "sync" + "sync/atomic" + "time" + + "github.com/google/uuid" + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/application" + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/trace" + grpcpkg "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + "github.com/mudler/xlog" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" + "google.golang.org/protobuf/proto" +) + +const motionIdleTimeout = 2 * time.Minute + +type MotionSessionRequest struct { + Model string `json:"model"` + Profile string `json:"profile"` + TimeOrigin string `json:"time_origin"` +} +type MotionSessionInfo struct { + UploadAccepted uint64 `json:"upload_accepted"` + UploadDropped uint64 `json:"upload_dropped"` + ID string `json:"id"` + Model string `json:"model"` + Profile string `json:"profile"` + TimeOrigin string `json:"time_origin"` + Frames uint64 `json:"frames"` + Poses uint64 `json:"poses"` +} + +type motionSession struct { + uploadAccepted, uploadDropped atomic.Uint64 + id, owner, model, engine, profile, origin string + ctx context.Context + cancel context.CancelFunc + stream grpcpkg.MotionStreamClient + definition []byte + frameMu sync.Mutex + frames, poses atomic.Uint64 + last atomic.Int64 + once sync.Once + started time.Time + traceID string +} +type MotionEndpoints struct { + ingress metric.Int64Counter + app *application.Application + corsEnabled bool + mu sync.Mutex + sessions map[string]*motionSession + reserved map[string]bool + active metric.Int64UpDownCounter + outcomes, resets metric.Int64Counter + latency, stages metric.Float64Histogram +} + +func NewMotionEndpoints(app *application.Application) *MotionEndpoints { + meter := otel.Meter("github.com/mudler/LocalAI/motion") + m := &MotionEndpoints{app: app, sessions: map[string]*motionSession{}, reserved: map[string]bool{}} + if app != nil { + m.corsEnabled = app.ApplicationConfig().CORS && app.ApplicationConfig().CORSAllowOrigins != "" + } + m.active, _ = meter.Int64UpDownCounter("localai_motion_sessions", metric.WithDescription("Active motion sessions")) + m.ingress, _ = meter.Int64Counter("localai_motion_uploads", metric.WithDescription("Duplex frames accepted or dropped before inference")) + m.outcomes, _ = meter.Int64Counter("localai_motion_frames", metric.WithDescription("Motion outcomes and rejected input")) + m.resets, _ = meter.Int64Counter("localai_motion_resets") + m.latency, _ = meter.Float64Histogram("localai_motion_frame_duration_seconds", metric.WithUnit("s"), metric.WithExplicitBucketBoundaries(.001, .005, .01, .025, .05, .1, .25, .5, 1, 2.5, 5, 10, 30, 60, 120)) + m.stages, _ = meter.Float64Histogram("localai_motion_stage_duration_seconds", metric.WithUnit("s"), metric.WithExplicitBucketBoundaries(.001, .005, .01, .025, .05, .1, .25, .5, 1, 2.5, 5, 10, 30, 60, 120)) + return m +} +func motionOwner(c echo.Context) string { + if u := auth.GetUser(c); u != nil { + return u.ID + } + return "unauthenticated" +} +func (m *MotionEndpoints) allowed(c echo.Context, name string) bool { + u := auth.GetUser(c) + return u == nil || auth.IsModelAllowed(m.app.AuthDB(), u, name) +} +func (m *MotionEndpoints) attrs(s *motionSession) metric.MeasurementOption { + return metric.WithAttributes(attribute.String("model", s.model), attribute.String("backend", s.engine)) +} +func (s *motionSession) info() MotionSessionInfo { + return MotionSessionInfo{UploadAccepted: s.uploadAccepted.Load(), UploadDropped: s.uploadDropped.Load(), ID: s.id, Model: s.model, Profile: s.profile, TimeOrigin: s.origin, Frames: s.frames.Load(), Poses: s.poses.Load()} +} +func (m *MotionEndpoints) lookup(c echo.Context) (*motionSession, error) { + m.mu.Lock() + s := m.sessions[c.Param("id")] + m.mu.Unlock() + if s == nil || s.owner != motionOwner(c) { + return nil, echo.NewHTTPError(404, "motion session not found") + } + if !m.allowed(c, s.model) { + return nil, echo.NewHTTPError(403, "model not allowed") + } + return s, nil +} +func (m *MotionEndpoints) close(s *motionSession, reason string) { + s.once.Do(func() { + s.cancel() + if m.app != nil { + m.app.WebSocketTickets().RevokePath("/api/motion/sessions/" + s.id + "/poses") + } + m.mu.Lock() + delete(m.sessions, s.id) + delete(m.reserved, s.model) + m.mu.Unlock() + m.active.Add(context.Background(), -1, m.attrs(s)) + failure := "" + switch reason { + case "backend error", "inference timeout", "invalid backend output": + failure = reason + } + if s.traceID != "" { + trace.RecordBackendTrace(trace.BackendTrace{ID: s.traceID, Timestamp: s.started, Duration: time.Since(s.started), Type: trace.BackendTraceMotion, ModelName: s.model, Backend: s.engine, Summary: "motion: " + reason, Error: failure, Data: map[string]any{"upload_accepted": s.uploadAccepted.Load(), "upload_dropped": s.uploadDropped.Load(), "frames": s.frames.Load(), "poses": s.poses.Load(), "profile": s.profile}}) + } + xlog.Info("motion session closed", "model", s.model, "reason", reason, "frames", s.frames.Load(), "poses", s.poses.Load(), "upload_accepted", s.uploadAccepted.Load(), "upload_dropped", s.uploadDropped.Load()) + }) +} + +// Create starts a resident motion session. +// @Summary Create a motion capture session +// @Tags motion +// @Accept json +// @Produce json +// @Param request body MotionSessionRequest true "Model, output profile and source clock origin" +// @Success 201 {object} MotionSessionInfo +// @Router /api/motion/sessions [post] +func (m *MotionEndpoints) Create(c echo.Context) error { + var req MotionSessionRequest + c.Request().Body = http.MaxBytesReader(c.Response(), c.Request().Body, 4096) + if err := c.Bind(&req); err != nil { + return echo.NewHTTPError(400, "invalid motion session request") + } + if req.Profile == "" { + req.Profile = "soma77" + } + if req.Model == "" || (req.Profile != "soma77" && req.Profile != "smpl24") || len(req.TimeOrigin) == 0 || len(req.TimeOrigin) > 128 { + return echo.NewHTTPError(400, "model, profile soma77/smpl24 and time_origin (1..128 characters) required") + } + if !m.allowed(c, req.Model) { + return echo.NewHTTPError(403, "model not allowed") + } + cfg, err := m.app.ModelConfigLoader().LoadModelConfigFileByNameDefaultOptions(req.Model, m.app.ApplicationConfig()) + if err != nil || cfg == nil { + return echo.NewHTTPError(404, "model not found") + } + if !cfg.HasUsecases(config.FLAG_MOTION) { + return echo.NewHTTPError(400, "model does not support motion") + } + m.mu.Lock() + if len(m.reserved) >= 16 || m.reserved[req.Model] { + m.mu.Unlock() + return echo.NewHTTPError(429, "motion capacity exhausted; one session per loaded model") + } + m.reserved[req.Model] = true + m.mu.Unlock() + success := false + defer func() { + if !success { + m.mu.Lock() + delete(m.reserved, req.Model) + m.mu.Unlock() + } + }() + ctx, cancel := context.WithCancel(m.app.ApplicationConfig().Context) + defer func() { + if !success { + cancel() + } + }() + stop := context.AfterFunc(c.Request().Context(), cancel) + defer stop() + loadTimer := time.AfterFunc(5*time.Minute, cancel) + defer loadTimer.Stop() + stream, err := backend.ModelMotionStream(ctx, m.app.ModelLoader(), m.app.ApplicationConfig(), *cfg) + if err != nil { + return echo.NewHTTPError(503, "motion backend unavailable") + } + if err := stream.Send(&pb.MotionRequest{ModelIdentity: cfg.Model, Profile: req.Profile}); err != nil { + return echo.NewHTTPError(502, "motion configuration failed") + } + ready, err := stream.Recv() + if err != nil { + return echo.NewHTTPError(502, "motion backend initialization failed") + } + out := &motion.Output{} + if err := proto.Unmarshal(ready.Output, out); err != nil || out.GetDefinition() == nil { + return echo.NewHTTPError(502, "invalid motion definition") + } + if out.GetDefinition().Conventions == nil { + out.GetDefinition().Conventions = map[string]string{} + } + out.GetDefinition().Conventions["time_origin"] = req.TimeOrigin + definition, err := proto.Marshal(out) + if err != nil { + return err + } + s := &motionSession{id: uuid.NewString(), owner: motionOwner(c), model: req.Model, engine: cfg.Backend, profile: req.Profile, origin: req.TimeOrigin, ctx: ctx, cancel: cancel, stream: stream, definition: definition, started: time.Now()} + s.last.Store(time.Now().UnixNano()) + if conf := m.app.ApplicationConfig(); conf.EnableTracing { + trace.InitBackendTracingIfEnabled(conf.TracingMaxItems, conf.TracingMaxBodyBytes) + s.traceID = trace.BeginBackendTrace(trace.BackendTrace{Timestamp: s.started, Type: trace.BackendTraceMotion, ModelName: s.model, Backend: s.engine, Summary: "motion: " + s.profile}) + } + m.mu.Lock() + m.sessions[s.id] = s + m.mu.Unlock() + success = true + m.active.Add(ctx, 1, m.attrs(s)) + xlog.Info("motion session started", "model", s.model, "profile", s.profile) + go func() { + ticker := time.NewTicker(10 * time.Second) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + m.close(s, "closed") + return + case <-ticker.C: + if time.Since(time.Unix(0, s.last.Load())) > motionIdleTimeout { + m.close(s, "idle timeout") + return + } + } + } + }() + return c.JSON(http.StatusCreated, s.info()) +} + +// Get returns session counters. +// @Summary Get motion session status +// @Tags motion +// @Param id path string true "Session ID" +// @Success 200 {object} MotionSessionInfo +// @Router /api/motion/sessions/{id} [get] +func (m *MotionEndpoints) Get(c echo.Context) error { + s, err := m.lookup(c) + if err != nil { + return err + } + return c.JSON(200, s.info()) +} + +// Delete releases the session. +// @Summary Close a motion session +// @Tags motion +// @Param id path string true "Session ID" +// @Success 204 +// @Router /api/motion/sessions/{id} [delete] +func (m *MotionEndpoints) Delete(c echo.Context) error { + s, err := m.lookup(c) + if err != nil { + return err + } + m.close(s, "deleted") + return c.NoContent(204) +} + +type motionFrameResult struct { + data []byte + output *motion.Output + elapsed time.Duration +} + +func (m *MotionEndpoints) processMotionInput(s *motionSession, input *motion.Input, data []byte) (*motionFrameResult, error) { + s.last.Store(time.Now().UnixNano()) + start := time.Now() + timeout := time.AfterFunc(motionIdleTimeout, func() { m.close(s, "inference timeout") }) + defer timeout.Stop() + if err := s.stream.Send(&pb.MotionRequest{Input: data}); err != nil { + m.close(s, "backend error") + return nil, echo.NewHTTPError(502, "motion backend send failed") + } + res, err := s.stream.Recv() + elapsed := time.Since(start) + if err != nil { + m.close(s, "backend error") + return nil, echo.NewHTTPError(502, "motion inference failed") + } + s.last.Store(time.Now().UnixNano()) + out := &motion.Output{} + if err := proto.Unmarshal(res.Output, out); err != nil { + m.close(s, "invalid backend output") + return nil, echo.NewHTTPError(502, "invalid motion output") + } + outcome := "pose" + flags := uint32(0) + if p := out.GetPose(); p != nil { + if frame := input.GetFrame(); frame != nil && (p.Sequence != frame.Sequence || p.SourceTimeUs != frame.SourceTimeUs) { + m.close(s, "backend pose identity mismatch") + return nil, echo.NewHTTPError(502, "backend pose identity mismatch") + } + s.poses.Add(1) + flags = p.Flags + } else if e := out.GetEvent(); e != nil { + switch e.Type { + case "warmup", "lost", "ambiguous", "reset": + outcome = e.Type + default: + m.close(s, "invalid backend output") + return nil, echo.NewHTTPError(502, "unexpected motion event") + } + flags = e.Flags + } else { + m.close(s, "invalid backend output") + return nil, echo.NewHTTPError(502, "unexpected motion output") + } + if f := input.GetFrame(); f != nil { + s.frames.Add(1) + } + m.outcomes.Add(s.ctx, 1, metric.WithAttributes(attribute.String("model", s.model), attribute.String("outcome", outcome))) + if flags&1 != 0 || outcome == "reset" { + m.resets.Add(s.ctx, 1, m.attrs(s)) + } + m.latency.Record(s.ctx, time.Since(start).Seconds(), m.attrs(s)) + for _, stage := range []string{"detector", "vitpose", "gem", "world", "camera"} { + if ms, ok := res.StageMs[stage]; ok && ms >= 0 && !math.IsNaN(ms) && !math.IsInf(ms, 0) { + m.stages.Record(s.ctx, ms/1000, metric.WithAttributes(attribute.String("model", s.model), attribute.String("stage", stage))) + } + } + return &motionFrameResult{data: res.Output, output: out, elapsed: elapsed}, nil +} + +type MotionTicketRequest struct { + Origin string `json:"origin"` +} + +// Ticket issues a single-use browser WebSocket credential for this session. +// @Summary Issue a motion WebSocket ticket (30 seconds, single use) +// @Tags motion +// @Accept json +// @Produce json +// @Param id path string true "Session ID" +// @Param request body MotionTicketRequest true "Browser origin" +// @Success 201 {object} auth.WebSocketTicketResponse +// @Router /api/motion/sessions/{id}/tickets [post] +func (m *MotionEndpoints) Ticket(c echo.Context) error { + s, err := m.lookup(c) + if err != nil { + return err + } + if auth.GetUser(c) == nil && (m.app.AuthDB() != nil || len(m.app.ApplicationConfig().ApiKeys) > 0) { + return echo.NewHTTPError(401, "authentication required") + } + c.Request().Body = http.MaxBytesReader(c.Response(), c.Request().Body, 4096) + var req MotionTicketRequest + if err := c.Bind(&req); err != nil { + return echo.NewHTTPError(400, "invalid ticket request") + } + ticket, err := m.app.WebSocketTickets().Issue(c, "/api/motion/sessions/"+s.id+"/poses", req.Origin) + if err != nil { + return err + } + c.Response().Header().Set("Cache-Control", "no-store") + return c.JSON(201, ticket) +} + +// Poses exchanges binary Protobuf frames and outputs; definition is always first. +// @Summary Stream motion frames and poses over a duplex WebSocket +// @Tags motion +// @Param id path string true "Session ID" +// @Success 101 +// @Router /api/motion/sessions/{id}/poses [get] +func (m *MotionEndpoints) Poses(c echo.Context) error { + s, err := m.lookup(c) + if err != nil { + return err + } + // Validated header credentials or an explicit CORS policy extend the origin default. + // Reuse middleware approval to keep origin matching consistent with HTTP. + duplex := false + for _, protocol := range websocket.Subprotocols(c.Request()) { + if protocol == "localai.motion.v2" { + duplex = true + } + } + if !duplex { + return echo.NewHTTPError(400, "localai.motion.v2 subprotocol required") + } + if !s.frameMu.TryLock() { + return echo.NewHTTPError(409, "motion session already has an input owner") + } + defer s.frameMu.Unlock() + upgrader := websocket.Upgrader{Subprotocols: []string{"localai.motion.v2"}} + origin := c.Request().Header.Get(echo.HeaderOrigin) + allowedOrigin := c.Response().Header().Get(echo.HeaderAccessControlAllowOrigin) + if auth.WebSocketTicketAuthenticated(c) || auth.HeaderAuthenticated(c) || (m.corsEnabled && origin != "" && (allowedOrigin == "*" || allowedOrigin == origin)) { + upgrader.CheckOrigin = func(*http.Request) bool { return true } + } + ws, err := upgrader.Upgrade(c.Response(), c.Request(), nil) + if err != nil { + return err + } + defer func() { _ = ws.Close() }() + return m.duplexMotion(c, s, ws) +} diff --git a/core/http/endpoints/localai/motion_duplex.go b/core/http/endpoints/localai/motion_duplex.go new file mode 100644 index 000000000000..86b5c1b29013 --- /dev/null +++ b/core/http/endpoints/localai/motion_duplex.go @@ -0,0 +1,357 @@ +// SPDX-License-Identifier: MIT +package localai + +import ( + "context" + "encoding/json" + "errors" + "strconv" + "sync" + "time" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/services/routing/billing" + "github.com/mudler/LocalAI/pkg/grpc/metadata" + motionutil "github.com/mudler/LocalAI/pkg/motion" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + "github.com/mudler/xlog" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" + "google.golang.org/protobuf/proto" +) + +// Each release retires an outstanding capture exactly once. Binary output precedes +// its completion release; thinning/expiry releases have no binary output. +type motionFlow struct { + Type string `json:"type"` + Version int `json:"version"` + ID uint64 `json:"id"` + Release []string `json:"release"` + CompletedSequence string `json:"completed_sequence,omitempty"` + CompletedKind string `json:"completed_kind,omitempty"` + WindowFrames int `json:"window_frames"` + MaxPendingBytes int `json:"max_pending_bytes"` + MaxFrameBytes int `json:"max_frame_bytes"` + ProcessingEWMA float64 `json:"processing_ewma_ms"` + RecommendedFPS float64 `json:"recommended_fps"` + QueuedFrames int `json:"queued_frames"` + QueuedBytes int `json:"queued_bytes"` + OldestAgeMS float64 `json:"oldest_age_ms"` + Pressure bool `json:"pressure"` + Dropped uint64 `json:"dropped_frames"` +} + +type motionSocketMessage struct { + output []byte + flow motionFlow +} + +// The reader admits monotonically ordered inputs; one worker owns inference. +// A single writer consumes this bounded channel, preserving output/release order. +type motionUpload struct { + recordIngress func(string, int) + mu sync.Mutex + pending motionutil.FrameQueue + lastSequence uint64 + lastTime int64 + haveFrame bool + processing time.Duration + pressure bool + wake chan struct{} + messages chan motionSocketMessage +} + +func newMotionUpload() *motionUpload { + return &motionUpload{ + recordIngress: func(string, int) {}, + wake: make(chan struct{}, 1), + messages: make(chan motionSocketMessage, 8), + } +} + +func (u *motionUpload) feedback(released []uint64) motionFlow { + releases := make([]string, len(released)) + for i, sequence := range released { + releases[i] = strconv.FormatUint(sequence, 10) + } + flow := motionFlow{ + Type: "flow", + Version: 2, + Release: releases, + WindowFrames: motionutil.UploadWindow, + MaxPendingBytes: motionutil.UploadMaxBytes, + MaxFrameBytes: motionutil.UploadFrameMaxBytes, + ProcessingEWMA: u.processing.Seconds() * 1000, + QueuedFrames: u.pending.Len(), + QueuedBytes: u.pending.Bytes(), + OldestAgeMS: u.pending.OldestAge(time.Now()).Seconds() * 1000, + Pressure: u.pressure, + Dropped: u.pending.Dropped, + } + if u.processing == 0 { + flow.WindowFrames = 2 + } else { + factor := 1.05 + if u.pressure { + factor = .9 + } + flow.RecommendedFPS = factor / u.processing.Seconds() + } + return flow +} + +func (u *motionUpload) send(ctx context.Context, msg motionSocketMessage) bool { + select { + case u.messages <- msg: + return true + case <-ctx.Done(): + return false + } +} + +func (u *motionUpload) admit(ctx context.Context, data []byte) error { + if err := ctx.Err(); err != nil { + return err + } + input := &motion.Input{} + if err := proto.Unmarshal(data, input); err != nil { + return errors.New("invalid motion protobuf") + } + frame := input.GetFrame() + if input.GetResetState() || frame == nil { + return errors.New("v2 accepts frames only; create a new session to reset") + } + if err := motionutil.ValidateFrame(frame); err != nil { + return err + } + + u.mu.Lock() + defer u.mu.Unlock() + if u.haveFrame && (frame.Sequence <= u.lastSequence || frame.SourceTimeUs <= u.lastTime) { + return errors.New("sequence and source time must increase") + } + u.lastSequence, u.lastTime, u.haveFrame = frame.Sequence, frame.SourceTimeUs, true + arrived := time.Now() + expired := u.pending.Expire(arrived) + u.recordIngress("expired", len(expired)) + released, pressure := u.pending.Offer(motionutil.QueuedFrame{Sequence: frame.Sequence, SourceTime: frame.SourceTimeUs, Data: data, Arrived: arrived}) + u.recordIngress("accepted", 1) + u.recordIngress("thinned", len(released)) + released = append(expired, released...) + u.pressure = u.pressure || pressure + if !u.send(ctx, motionSocketMessage{flow: u.feedback(released)}) { + return ctx.Err() + } + return nil +} + +func (u *motionUpload) read(s *motionSession, ws *websocket.Conn) error { + ws.SetReadLimit(motionutil.UploadFrameMaxBytes) + for { + kind, data, err := ws.ReadMessage() + if err != nil { + return err + } + if kind != websocket.BinaryMessage { + return errors.New("binary motion Input required") + } + err = u.admit(s.ctx, data) + if err != nil { + return err + } + s.last.Store(time.Now().UnixNano()) + select { + case u.wake <- struct{}{}: + default: + } + } +} + +// Exactly one worker holds the session's frameMu through the v2 owner lifetime. +func (m *MotionEndpoints) runMotionUpload(s *motionSession, u *motionUpload, record func(*motionFrameResult)) error { + for { + u.mu.Lock() + expired := u.pending.Expire(time.Now()) + u.recordIngress("expired", len(expired)) + frame, ok := u.pending.Take() + if u.pending.Len() <= 2 { + u.pressure = false + } + flow := u.feedback(expired) + if len(expired) > 0 && !u.send(s.ctx, motionSocketMessage{flow: flow}) { + u.mu.Unlock() + return s.ctx.Err() + } + u.mu.Unlock() + if !ok { + select { + case <-u.wake: + continue + case <-s.ctx.Done(): + return s.ctx.Err() + } + } + if s.ctx.Err() != nil { + return s.ctx.Err() + } + input := &motion.Input{} + if err := proto.Unmarshal(frame.Data, input); err != nil { + return err + } + result, err := m.processMotionInput(s, input, frame.Data) + if err != nil { + return err + } + record(result) + + kind := "event" + if pose := result.output.GetPose(); pose != nil { + kind = "pose" + } + u.mu.Lock() + elapsed := max(time.Millisecond, result.elapsed) + if u.processing == 0 { + u.processing = elapsed + } else { + u.processing = (u.processing*4 + elapsed) / 5 + } + flow = u.feedback([]uint64{frame.Sequence}) + flow.CompletedSequence = strconv.FormatUint(frame.Sequence, 10) + flow.CompletedKind = kind + sent := u.send(s.ctx, motionSocketMessage{output: result.data, flow: flow}) + u.mu.Unlock() + if !sent { + return s.ctx.Err() + } + } +} + +func writeMotionSocket(ws *websocket.Conn, kind int, data []byte) error { + _ = ws.SetWriteDeadline(time.Now().Add(5 * time.Second)) + return ws.WriteMessage(kind, data) +} + +func (m *MotionEndpoints) duplexMotion(c echo.Context, s *motionSession, ws *websocket.Conn) error { + u := newMotionUpload() + u.recordIngress = func(outcome string, count int) { + if count == 0 { + return + } + if outcome == "accepted" { + s.uploadAccepted.Add(uint64(count)) + } else { + s.uploadDropped.Add(uint64(count)) + } + m.ingress.Add(s.ctx, int64(count), metric.WithAttributes(attribute.String("model", s.model), attribute.String("outcome", outcome))) + } + record := m.motionSocketAccounting(c, s) + defer m.close(s, "duplex disconnected") + if err := writeMotionSocket(ws, websocket.BinaryMessage, s.definition); err != nil { + return nil + } + flowID := uint64(0) + writeFlow := func(flow motionFlow) error { + flowID++ + flow.ID = flowID + data, err := json.Marshal(flow) + if err != nil { + return err + } + return writeMotionSocket(ws, websocket.TextMessage, data) + } + if err := writeFlow(u.feedback(nil)); err != nil { + return nil + } + + _ = ws.SetReadDeadline(time.Now().Add(time.Minute)) + ws.SetPongHandler(func(string) error { return ws.SetReadDeadline(time.Now().Add(time.Minute)) }) + finished := make(chan error, 2) + var workers sync.WaitGroup + workers.Add(2) + go func() { defer workers.Done(); finished <- u.read(s, ws) }() + go func() { defer workers.Done(); finished <- m.runMotionUpload(s, u, record) }() + // Cancellation closes gRPC transport; it does not claim to preempt native GPU work. + defer func() { m.close(s, "duplex disconnected"); _ = ws.Close(); workers.Wait() }() + ping := time.NewTicker(20 * time.Second) + defer ping.Stop() + for { + select { + case err := <-finished: + if err != nil && s.ctx.Err() == nil { + xlog.Debug("motion duplex closed", "model", s.model, "error", err) + } + _ = ws.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "motion stream ended; create a new session"), time.Now().Add(time.Second)) + return nil + case <-s.ctx.Done(): + return nil + case msg := <-u.messages: + if msg.output != nil { + if err := writeMotionSocket(ws, websocket.BinaryMessage, msg.output); err != nil { + return nil + } + } + if err := writeFlow(msg.flow); err != nil { + return nil + } + case <-ping.C: + if err := ws.WriteControl(websocket.PingMessage, nil, time.Now().Add(5*time.Second)); err != nil { + return nil + } + } + } +} + +// WebSocket upgrades bypass HTTP usage middleware. Capture identity before worker +// startup and record processed inputs directly through the same billing recorder. +func (m *MotionEndpoints) motionSocketAccounting(c echo.Context, s *motionSession) func(*motionFrameResult) { + if m.app == nil { + return func(*motionFrameResult) {} + } + return recordMotionUsage(c, s, m.app.StatsRecorder(), m.app.FallbackUser()) +} + +func recordMotionUsage(c echo.Context, s *motionSession, recorder *billing.Recorder, fallback *auth.User) func(*motionFrameResult) { + if recorder == nil { + return func(*motionFrameResult) {} + } + user := auth.GetUser(c) + if user == nil { + user = fallback + } + endpoint := c.Request().URL.Path + if user == nil { + return func(*motionFrameResult) { billing.CountUnrecorded(context.Background(), endpoint, "no_user") } + } + base := auth.UsageRecord{UserID: user.ID, UserName: user.Name, Source: auth.GetSource(c), Model: s.model, Endpoint: endpoint, RequestedModel: s.model, ServedModel: s.model} + if base.Source == "" { + base.Source = auth.UsageSourceWeb + } + if key := auth.GetAPIKey(c); key != nil { + id := key.ID + base.APIKeyID = &id + } + return func(result *motionFrameResult) { + outputs := 0 + if result.output.GetPose() != nil { + outputs = 1 + } + usage, err := metadata.EncodeUsage(metadata.Usage{InputUnits: 1, OutputUnits: outputs, AccountingRule: "motion_frames_v1"}) + if err != nil { + xlog.Error("invalid motion usage metadata", "error", err) + billing.CountUnrecorded(context.Background(), endpoint, "invalid_usage") + return + } + record := base + record.PromptTokens, record.CompletionTokens, record.TotalTokens = 1, int64(outputs), int64(1+outputs) + record.PreFilterPromptTokens, record.PostFilterPromptTokens = 1, 1 + record.Duration = result.elapsed.Milliseconds() + record.CreatedAt = time.Now() + record.Metadata = string(usage) + if err := recorder.Record(context.Background(), &record); err != nil { + xlog.Error("motion usage recording failed", "error", err) + billing.CountUnrecorded(context.Background(), endpoint, "record_error") + } + } +} diff --git a/core/http/endpoints/localai/motion_duplex_test.go b/core/http/endpoints/localai/motion_duplex_test.go new file mode 100644 index 000000000000..2682a2307838 --- /dev/null +++ b/core/http/endpoints/localai/motion_duplex_test.go @@ -0,0 +1,246 @@ +// SPDX-License-Identifier: MIT +package localai + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "sync" + "time" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/protobuf/proto" +) + +type gatedMotionStream struct { + ctx context.Context + gate chan struct{} + started chan uint64 + input *motion.Input + mu sync.Mutex + processed []uint64 + wrongIdentity bool +} + +func (f *gatedMotionStream) Send(r *pb.MotionRequest) error { + f.input = &motion.Input{} + return proto.Unmarshal(r.Input, f.input) +} +func (f *gatedMotionStream) Recv() (*pb.MotionResponse, error) { + frame := f.input.GetFrame() + select { + case f.started <- frame.Sequence: + default: + } + select { + case <-f.gate: + case <-f.ctx.Done(): + return nil, f.ctx.Err() + } + f.mu.Lock() + f.processed = append(f.processed, frame.Sequence) + f.mu.Unlock() + out := &motion.Output{Payload: &motion.Output_Pose{Pose: &motion.Pose{Sequence: frame.Sequence, SourceTimeUs: frame.SourceTimeUs}}} + if f.wrongIdentity { + out.GetPose().Sequence++ + } + data, err := proto.Marshal(out) + return &pb.MotionResponse{Output: data}, err +} +func (f *gatedMotionStream) CloseSend() error { return nil } +func (f *gatedMotionStream) Context() context.Context { return f.ctx } + +func duplexTestFrame(sequence uint64) []byte { + data, err := proto.Marshal(&motion.Input{Payload: &motion.Input_Frame{Frame: &motion.Frame{Sequence: sequence, SourceTimeUs: int64(sequence) * 10000, Rgb: make([]byte, 8*8*3), Width: 8, Height: 8}}}) + Expect(err).NotTo(HaveOccurred()) + return data +} + +var _ = Describe("Motion duplex WebSocket", func() { + var m *MotionEndpoints + var s *motionSession + var backend *gatedMotionStream + var server *httptest.Server + var address string + BeforeEach(func() { + m = NewMotionEndpoints(nil) + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + backend = &gatedMotionStream{ + ctx: ctx, + gate: make(chan struct{}), + started: make(chan uint64, 8), + } + s = &motionSession{ + id: "duplex", + owner: "unauthenticated", + model: "gem-x", + profile: "smpl24", + ctx: ctx, + cancel: cancel, + stream: backend, + started: time.Now(), + } + s.definition = motionBytes(&motion.Output{Payload: &motion.Output_Definition{Definition: &motion.Definition{Schema: "test"}}}) + m.sessions[s.id] = s + m.reserved[s.model] = true + + e := echo.New() + e.GET("/api/motion/sessions/:id/poses", m.Poses) + server = httptest.NewServer(e) + DeferCleanup(server.Close) + address = "ws" + strings.TrimPrefix(server.URL, "http") + "/api/motion/sessions/duplex/poses" + }) + connect := func() *websocket.Conn { + + dialer := websocket.Dialer{Subprotocols: []string{"localai.motion.v2"}} + ws, response, err := dialer.Dial(address, nil) + if response != nil { + DeferCleanup(response.Body.Close) + } + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { _ = ws.Close() }) + + Expect(ws.SetReadDeadline(time.Now().Add(3 * time.Second))).To(Succeed()) + Expect(ws.Subprotocol()).To(Equal("localai.motion.v2")) + + kind, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + Expect(kind).To(Equal(websocket.BinaryMessage)) + Expect(data).To(Equal(s.definition)) + + kind, data, err = ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + Expect(kind).To(Equal(websocket.TextMessage)) + var flow motionFlow + Expect(json.Unmarshal(data, &flow)).To(Succeed()) + Expect(flow.ID).To(Equal(uint64(1))) + Expect(flow.WindowFrames).To(Equal(2)) + Expect(flow.RecommendedFPS).To(BeZero()) + return ws + } + DescribeTable("rejects non-duplex protocols before upgrade", func(protocols []string) { + d := websocket.Dialer{Subprotocols: protocols} + _, response, err := d.Dial(address, nil) + Expect(err).To(HaveOccurred()) + Expect(response.StatusCode).To(Equal(http.StatusBadRequest)) + Expect(response.Body.Close()).To(Succeed()) + Expect(s.ctx.Err()).NotTo(HaveOccurred()) + Expect(s.frameMu.TryLock()).To(BeTrue()) + s.frameMu.Unlock() + }, Entry("missing", []string{}), Entry("read-only", []string{"localai.motion.v1"})) + It("thins pending frames evenly, preserves active inference, and releases each ID once", func() { + ws := connect() + Expect(ws.WriteMessage(websocket.BinaryMessage, duplexTestFrame(1))).To(Succeed()) + Eventually(backend.started).Should(Receive(Equal(uint64(1)))) + + // Deliberately overload the server rather than obeying its bootstrap credits. + for sequence := uint64(2); sequence <= 7; sequence++ { + Expect(ws.WriteMessage(websocket.BinaryMessage, duplexTestFrame(sequence))).To(Succeed()) + } + + released := map[string]bool{} + sawPressure := false + for !sawPressure { + _, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + var flow motionFlow + Expect(json.Unmarshal(data, &flow)).To(Succeed()) + for _, id := range flow.Release { + Expect(released[id]).To(BeFalse()) + released[id] = true + } + if flow.Pressure { + Expect(flow.Release).To(Equal([]string{"3", "5", "6"})) + Expect(flow.QueuedFrames).To(Equal(3)) + sawPressure = true + } + } + + close(backend.gate) + for len(released) < 7 { + kind, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + if kind == websocket.BinaryMessage { + continue + } + var flow motionFlow + Expect(json.Unmarshal(data, &flow)).To(Succeed()) + for _, id := range flow.Release { + Expect(released[id]).To(BeFalse()) + released[id] = true + } + Expect(flow.RecommendedFPS).To(BeNumerically(">", 0)) + } + + backend.mu.Lock() + processed := append([]uint64(nil), backend.processed...) + backend.mu.Unlock() + Expect(processed).To(Equal([]uint64{1, 2, 4, 7})) + }) + It("rejects competing ingress and cancels a blocked worker on disconnect", func() { + ws := connect() + Expect(ws.WriteMessage(websocket.BinaryMessage, duplexTestFrame(1))).To(Succeed()) + Eventually(backend.started).Should(Receive()) + + dialer := websocket.Dialer{Subprotocols: []string{"localai.motion.v2"}} + _, response, err := dialer.Dial(address, nil) + Expect(err).To(HaveOccurred()) + Expect(response.StatusCode).To(Equal(http.StatusConflict)) + Expect(response.Body.Close()).To(Succeed()) + + Expect(ws.Close()).To(Succeed()) + Eventually(s.ctx.Done()).Should(BeClosed()) + Eventually(func() bool { + if s.frameMu.TryLock() { + s.frameMu.Unlock() + return true + } + return false + }).Should(BeTrue()) + + m.mu.Lock() + remaining := len(m.sessions) + m.mu.Unlock() + Expect(remaining).To(BeZero()) + }) + It("rejects reset without backend work", func() { + ws := connect() + Expect(ws.WriteMessage(websocket.BinaryMessage, []byte{16, 1})).To(Succeed()) + Eventually(s.ctx.Done()).Should(BeClosed()) + Consistently(backend.started).ShouldNot(Receive()) + }) + It("rejects oversized socket frames before protobuf decoding", func() { + ws := connect() + _ = ws.WriteMessage(websocket.BinaryMessage, make([]byte, (2<<20)+1)) + Eventually(s.ctx.Done()).Should(BeClosed()) + Consistently(backend.started).ShouldNot(Receive()) + }) + It("rejects sequence replay even when the original was dropped", func() { + u := newMotionUpload() + for sequence := uint64(1); sequence <= 6; sequence++ { + Expect(u.admit(s.ctx, duplexTestFrame(sequence))).To(Succeed()) + } + Expect(u.admit(s.ctx, duplexTestFrame(2))).To(MatchError("sequence and source time must increase")) + }) + It("rejects mismatched backend poses before publication or counters", func() { + backend.wrongIdentity = true + ws := connect() + + Expect(ws.WriteMessage(websocket.BinaryMessage, duplexTestFrame(1))).To(Succeed()) + Eventually(backend.started).Should(Receive()) + close(backend.gate) + + Eventually(s.ctx.Done()).Should(BeClosed()) + Expect(s.poses.Load()).To(BeZero()) + Expect(s.frames.Load()).To(BeZero()) + }) + +}) diff --git a/core/http/endpoints/localai/motion_internal_test.go b/core/http/endpoints/localai/motion_internal_test.go new file mode 100644 index 000000000000..845c22cbaccd --- /dev/null +++ b/core/http/endpoints/localai/motion_internal_test.go @@ -0,0 +1,323 @@ +// SPDX-License-Identifier: MIT +package localai + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "time" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + echomiddleware "github.com/labstack/echo/v4/middleware" + "github.com/mudler/LocalAI/core/application" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/services/routing/billing" + "github.com/mudler/LocalAI/pkg/grpc/metadata" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + motion "github.com/mudler/LocalAI/pkg/motion/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "go.opentelemetry.io/otel" + sdkmetric "go.opentelemetry.io/otel/sdk/metric" + "go.opentelemetry.io/otel/sdk/metric/metricdata" + "google.golang.org/protobuf/proto" +) + +type fakeMotionStream struct { + ctx context.Context + input *pb.MotionRequest + response *pb.MotionResponse + err error +} + +func (f *fakeMotionStream) Send(r *pb.MotionRequest) error { f.input = r; return nil } +func (f *fakeMotionStream) Recv() (*pb.MotionResponse, error) { return f.response, f.err } +func (f *fakeMotionStream) CloseSend() error { return nil } +func (f *fakeMotionStream) Context() context.Context { return f.ctx } +func motionBytes(o *motion.Output) []byte { + b, err := proto.Marshal(o) + Expect(err).NotTo(HaveOccurred()) + return b +} + +type motionUsageBackend struct { + billing.StatsBackend + records []*auth.UsageRecord +} + +func (b *motionUsageBackend) Record(ctx context.Context, r *auth.UsageRecord) error { + b.records = append(b.records, r) + return b.StatsBackend.Record(ctx, r) +} + +var configForMotionAuth = config.ApplicationConfig{ApiKeys: []string{"test-key"}} + +var _ = Describe("Motion streaming endpoints", func() { + var m *MotionEndpoints + var s *motionSession + var reader *sdkmetric.ManualReader + BeforeEach(func() { + reader = sdkmetric.NewManualReader() + provider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) + previous := otel.GetMeterProvider() + otel.SetMeterProvider(provider) + DeferCleanup(func() { otel.SetMeterProvider(previous); Expect(provider.Shutdown(context.Background())).To(Succeed()) }) + m = NewMotionEndpoints(nil) + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + s = &motionSession{id: "session", owner: "unauthenticated", model: "gem-x", engine: "gemxcpp", profile: "smpl24", ctx: ctx, cancel: cancel, started: time.Now()} + s.definition = motionBytes(&motion.Output{Payload: &motion.Output_Definition{Definition: &motion.Definition{Schema: "test", Profile: "smpl24"}}}) + m.sessions[s.id] = s + m.reserved[s.model] = true + }) + It("does not expose another user's session", func() { + c := echo.New().NewContext(httptest.NewRequest("GET", "/", nil), httptest.NewRecorder()) + c.SetParamNames("id") + c.SetParamValues(s.id) + c.Set("auth_user", &auth.User{ID: "different"}) + _, err := m.lookup(c) + Expect(err).To(MatchError(echo.NewHTTPError(404, "motion session not found"))) + }) + It("rejects unauthenticated access through standard middleware", func() { + e := echo.New() + e.Use(auth.Middleware(nil, &configForMotionAuth)) + e.GET("/api/motion/sessions/:id", m.Get) + recorder := httptest.NewRecorder() + e.ServeHTTP(recorder, httptest.NewRequest("GET", "/api/motion/sessions/session", nil)) + Expect(recorder.Code).To(Equal(http.StatusUnauthorized)) + }) + It("preserves source timestamps through ingestion and records stage metrics", func() { + frame := &motion.Frame{Width: 8, Height: 8, Rgb: make([]byte, 192), Sequence: 7, SourceTimeUs: 123456} + input, err := proto.Marshal(&motion.Input{Payload: &motion.Input_Frame{Frame: frame}}) + Expect(err).NotTo(HaveOccurred()) + output := motionBytes(&motion.Output{Payload: &motion.Output_Pose{Pose: &motion.Pose{Sequence: 7, SourceTimeUs: 123456, Flags: 1}}}) + stream := &fakeMotionStream{ctx: s.ctx, response: &pb.MotionResponse{Output: output, StageMs: map[string]float64{"gem": 12}}} + s.stream = stream + decoded := &motion.Input{} + Expect(proto.Unmarshal(input, decoded)).To(Succeed()) + result, err := m.processMotionInput(s, decoded, input) + Expect(err).NotTo(HaveOccurred()) + Expect(result.data).To(Equal(output)) + Expect(stream.input.Input).To(Equal(input)) + Expect(s.frames.Load()).To(Equal(uint64(1))) + Expect(s.poses.Load()).To(Equal(uint64(1))) + var metrics metricdata.ResourceMetrics + Expect(reader.Collect(context.Background(), &metrics)).To(Succeed()) + names := []string{} + for _, scope := range metrics.ScopeMetrics { + for _, metric := range scope.Metrics { + names = append(names, metric.Name) + } + } + Expect(names).To(ContainElements("localai_motion_frames", "localai_motion_resets", "localai_motion_stage_duration_seconds")) + }) + + DescribeTable("accounts once per processed socket frame", func(outcome string, wantOutput int, enabled bool) { + stats := &motionUsageBackend{StatsBackend: billing.NewMemoryBackend(10)} + DeferCleanup(func() { Expect(stats.Close()).To(Succeed()) }) + var recorder *billing.Recorder + if enabled { + recorder = billing.NewRecorder(stats) + } + c := echo.New().NewContext(httptest.NewRequest("GET", "/api/motion/sessions/session/poses", nil), httptest.NewRecorder()) + record := recordMotionUsage(c, s, recorder, &auth.User{ID: "local"}) + output := &motion.Output{Payload: &motion.Output_Event{Event: &motion.Event{Type: outcome}}} + if outcome == "pose" { + output.Payload = &motion.Output_Pose{Pose: &motion.Pose{}} + } + record(&motionFrameResult{output: output, elapsed: time.Millisecond}) + if !enabled { + Expect(stats.records).To(BeEmpty()) + return + } + Expect(stats.records).To(HaveLen(1)) + entry := stats.records[0] + Expect(entry.UserID).To(Equal("local")) + Expect(entry.Model).To(Equal("gem-x")) + Expect(entry.PromptTokens).To(Equal(int64(1))) + Expect(entry.CompletionTokens).To(Equal(int64(wantOutput))) + usage, err := metadata.ParseUsage([]byte(entry.Metadata)) + Expect(err).NotTo(HaveOccurred()) + Expect(usage.AccountingRule).To(Equal("motion_frames_v1")) + }, Entry("pose", "pose", 1, true), Entry("warmup", "warmup", 0, true), Entry("lost", "lost", 0, true), Entry("ambiguous", "ambiguous", 0, true), Entry("disabled", "pose", 1, false)) + + DescribeTable("rejects failed or malformed inference without counting a frame", func(failure string) { + stream := &fakeMotionStream{ctx: s.ctx} + if failure == "backend" { + stream.err = errors.New("backend failed") + } else { + stream.response = &pb.MotionResponse{Output: motionBytes(&motion.Output{Payload: &motion.Output_Definition{Definition: &motion.Definition{}}})} + } + s.stream = stream + input := &motion.Input{} + data := duplexTestFrame(1) + Expect(proto.Unmarshal(data, input)).To(Succeed()) + _, err := m.processMotionInput(s, input, data) + Expect(err).To(HaveOccurred()) + Expect(s.frames.Load()).To(BeZero()) + Expect(s.poses.Load()).To(BeZero()) + Expect(s.ctx.Done()).To(BeClosed()) + }, Entry("backend error", "backend"), Entry("invalid output", "malformed")) + + It("rejects cross-origin browser connections", func() { + e := echo.New() + e.GET("/api/motion/sessions/:id/poses", m.Poses) + server := httptest.NewServer(e) + DeferCleanup(server.Close) + _, response, err := (&websocket.Dialer{Subprotocols: []string{"localai.motion.v2"}}).Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/api/motion/sessions/session/poses", http.Header{"Origin": []string{"https://untrusted.example"}}) + Expect(err).To(HaveOccurred()) + Expect(response.StatusCode).To(Equal(403)) + Expect(response.Body.Close()).To(Succeed()) + }) + DescribeTable("honors explicitly enabled CORS for browser connections", func(enabled bool, allowed []string, origin string, accept bool) { + m.corsEnabled = enabled + e := echo.New() + e.Use(echomiddleware.CORSWithConfig(echomiddleware.CORSConfig{AllowOrigins: allowed})) + e.GET("/api/motion/sessions/:id/poses", m.Poses) + server := httptest.NewServer(e) + DeferCleanup(server.Close) + ws, response, err := (&websocket.Dialer{Subprotocols: []string{"localai.motion.v2"}}).Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/api/motion/sessions/session/poses", http.Header{"Origin": []string{origin}}) + if !accept { + Expect(err).To(HaveOccurred()) + Expect(response.StatusCode).To(Equal(403)) + Expect(response.Body.Close()).To(Succeed()) + return + } + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { Expect(ws.Close()).To(Succeed()) }) + Expect(response.StatusCode).To(Equal(101)) + kind, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + Expect(kind).To(Equal(websocket.BinaryMessage)) + out := &motion.Output{} + Expect(proto.Unmarshal(data, out)).To(Succeed()) + Expect(out.GetDefinition().Profile).To(Equal("smpl24")) + }, + Entry("explicit allowed origin", true, []string{"https://consumer.example"}, "https://consumer.example", true), + Entry("explicit wildcard policy", true, []string{"*"}, "http://localhost:3000", true), + Entry("subdomain policy", true, []string{"https://*.example.com"}, "https://consumer.example.com", true), + Entry("disallowed origin", true, []string{"https://consumer.example"}, "https://untrusted.example", false), + Entry("different port", true, []string{"http://localhost:3000"}, "http://localhost:3001", false), + Entry("default HTTP wildcard does not enable cross-origin sockets", false, []string{"*"}, "https://untrusted.example", false), + ) + + DescribeTable("bypasses origins only for validated header credentials", func(header, value string, status int) { + m.app = &application.Application{} + s.owner = "legacy-api-key" + e := echo.New() + e.Use(auth.Middleware(nil, &configForMotionAuth)) + e.GET("/api/motion/sessions/:id/poses", m.Poses) + server := httptest.NewServer(e) + DeferCleanup(server.Close) + headers := http.Header{"Origin": []string{"https://external.example"}} + headers.Set(header, value) + ws, response, err := (&websocket.Dialer{Subprotocols: []string{"localai.motion.v2"}}).Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/api/motion/sessions/session/poses", headers) + Expect(response.StatusCode).To(Equal(status)) + if status != 101 { + Expect(err).To(HaveOccurred()) + Expect(response.Body.Close()).To(Succeed()) + return + } + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { Expect(ws.Close()).To(Succeed()) }) + _, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + output := &motion.Output{} + Expect(proto.Unmarshal(data, output)).To(Succeed()) + Expect(output.GetDefinition()).NotTo(BeNil()) + }, + Entry("Bearer", "Authorization", "Bearer test-key", 101), + Entry("API key", "x-api-key", "test-key", 101), + Entry("alternative API key", "xi-api-key", "test-key", 101), + Entry("invalid key", "Authorization", "Bearer invalid", 401), + Entry("cookie key retains origin policy", "Cookie", "token=test-key", 403), + ) + + It("issues browser tickets, preserves identity and rejects replay and revoked credentials", func() { + m.app = &application.Application{} + s.owner = "legacy-api-key" + cfg := config.ApplicationConfig{ApiKeys: []string{"test-key"}} + e := echo.New() + e.Use(auth.WithWebSocketTickets(m.app.WebSocketTickets(), auth.Middleware(nil, &cfg))) + e.POST("/api/motion/sessions/:id/tickets", m.Ticket) + e.GET("/api/motion/sessions/:id/poses", m.Poses) + server := httptest.NewServer(e) + DeferCleanup(server.Close) + issue := func() auth.WebSocketTicketResponse { + req := httptest.NewRequest("POST", "/api/motion/sessions/session/tickets", strings.NewReader(`{"origin":"https://consumer.example"}`)) + req.Header.Set("Authorization", "Bearer test-key") + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", "https://consumer.example") + res := httptest.NewRecorder() + e.ServeHTTP(res, req) + Expect(res.Code).To(Equal(201), res.Body.String()) + Expect(res.Header().Get("Cache-Control")).To(Equal("no-store")) + var ticket auth.WebSocketTicketResponse + Expect(json.Unmarshal(res.Body.Bytes(), &ticket)).To(Succeed()) + return ticket + } + connect := func(ticket, origin, id string) (*websocket.Conn, *http.Response, error) { + d := websocket.Dialer{Subprotocols: []string{"localai.motion.v2", auth.WebSocketTicketProtocolPrefix + ticket}} + return d.Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/api/motion/sessions/"+id+"/poses", http.Header{"Origin": []string{origin}}) + } + ticket := issue() + revoked := issue() + for _, attempt := range []struct{ origin, id string }{{"https://untrusted.example", "session"}, {"https://consumer.example", "other"}} { + _, res, err := connect(ticket.Ticket, attempt.origin, attempt.id) + Expect(err).To(HaveOccurred()) + Expect(res.StatusCode).To(Equal(401)) + Expect(res.Body.Close()).To(Succeed()) + } + ws, res, err := connect(ticket.Ticket, "https://consumer.example", "session") + Expect(err).NotTo(HaveOccurred()) + Expect(ws.Subprotocol()).To(Equal("localai.motion.v2")) + Expect(res.Header.Get("Sec-WebSocket-Protocol")).NotTo(ContainSubstring(ticket.Ticket)) + _, data, err := ws.ReadMessage() + Expect(err).NotTo(HaveOccurred()) + out := &motion.Output{} + Expect(proto.Unmarshal(data, out)).To(Succeed()) + Expect(out.GetDefinition()).NotTo(BeNil()) + Expect(ws.Close()).To(Succeed()) + _, res, err = connect(ticket.Ticket, "https://consumer.example", "session") + Expect(err).To(HaveOccurred()) + Expect(res.StatusCode).To(Equal(401)) + Expect(res.Body.Close()).To(Succeed()) + cfg.ApiKeys = []string{"replacement"} + _, res, err = connect(revoked.Ticket, "https://consumer.example", "session") + Expect(err).To(HaveOccurred()) + Expect(res.StatusCode).To(Equal(401)) + Expect(res.Body.Close()).To(Succeed()) + }) + It("does not issue tickets without session ownership or valid authentication", func() { + m.app = &application.Application{} + s.owner = "someone-else" + e := echo.New() + e.Use(auth.Middleware(nil, &configForMotionAuth)) + e.POST("/api/motion/sessions/:id/tickets", m.Ticket) + for _, key := range []string{"", "invalid", "test-key"} { + req := httptest.NewRequest("POST", "/api/motion/sessions/session/tickets", strings.NewReader(`{"origin":"https://consumer.example"}`)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+key) + res := httptest.NewRecorder() + e.ServeHTTP(res, req) + if key == "test-key" { + Expect(res.Code).To(Equal(404)) + } else { + Expect(res.Code).To(Equal(401)) + } + } + }) + + It("cancels once and releases the model reservation on deletion", func() { + m.close(s, "deleted") + m.close(s, "deleted") + Expect(s.ctx.Done()).To(BeClosed()) + Expect(m.sessions).To(BeEmpty()) + Expect(m.reserved).To(BeEmpty()) + }) +}) diff --git a/core/http/middleware/csrf_test.go b/core/http/middleware/csrf_test.go new file mode 100644 index 000000000000..ec3b828f3ab3 --- /dev/null +++ b/core/http/middleware/csrf_test.go @@ -0,0 +1,51 @@ +package middleware_test + +import ( + "net/http/httptest" + "strings" + + "github.com/labstack/echo/v4" + echoMiddleware "github.com/labstack/echo/v4/middleware" + "github.com/mudler/LocalAI/core/http/auth" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("CSRF trust of explicit CORS origins", func() { + DescribeTable("uses the HTTP origin decision for state-changing requests", func(explicit bool, allowlist, origin string, accepted bool) { + for _, method := range []string{"POST", "PUT", "PATCH", "DELETE"} { + e := echo.New() + if allowlist != "" { + e.Use(echoMiddleware.CORSWithConfig(echoMiddleware.CORSConfig{AllowOrigins: strings.Split(allowlist, ",")})) + } + e.Use(auth.CSRFMiddlewareWithCORS(explicit)) + called := false + e.Add(method, "/test", func(c echo.Context) error { called = true; return c.NoContent(204) }) + r := httptest.NewRequest(method, "/test", nil) + r.Header.Set("Origin", origin) + r.Header.Set("Sec-Fetch-Site", "cross-site") + // A caller cannot fabricate middleware approval with a request header. + r.Header.Set("Access-Control-Allow-Origin", "*") + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + Expect(called).To(Equal(accepted), method) + if accepted { + Expect(w.Code).To(Equal(204)) + } else { + Expect(w.Code).To(BeNumerically(">=", 400)) + Expect(w.Code).To(BeNumerically("<", 500)) + } + } + }, + Entry("allowed browser origin", true, "https://consumer.example", "https://consumer.example", true), + Entry("multiple origins", true, "https://first.example,https://consumer.example", "https://consumer.example", true), + Entry("disallowed origin", true, "https://consumer.example", "https://untrusted.example", false), + Entry("different port", true, "https://consumer.example", "https://consumer.example:8443", false), + Entry("different scheme", true, "https://consumer.example", "http://consumer.example", false), + Entry("explicit wildcard", true, "*", "https://consumer.example", true), + Entry("subdomain wildcard", true, "https://*.example.com", "https://app.example.com", true), + Entry("default permissive CORS does not disable CSRF", false, "*", "https://consumer.example", false), + Entry("empty strict allowlist", false, "", "https://consumer.example", false), + Entry("missing origin", true, "*", "", false), + ) +}) diff --git a/core/http/middleware/trace.go b/core/http/middleware/trace.go index 02d2e6437e4d..caa8af53bdf5 100644 --- a/core/http/middleware/trace.go +++ b/core/http/middleware/trace.go @@ -217,13 +217,14 @@ func (w *bodyWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { // land in the in-memory trace buffer. Keys are canonical — http.Header // stores them that way, so range yields canonical keys directly. var sensitiveTraceHeaders = map[string]struct{}{ - "Authorization": {}, - "Proxy-Authorization": {}, - "Cookie": {}, - "Set-Cookie": {}, - "X-Api-Key": {}, - "Xi-Api-Key": {}, - "X-Auth-Token": {}, + "Authorization": {}, + "Proxy-Authorization": {}, + "Cookie": {}, + "Set-Cookie": {}, + "X-Api-Key": {}, + "Xi-Api-Key": {}, + "X-Auth-Token": {}, + "Sec-Websocket-Protocol": {}, } func redactSensitiveHeaders(h http.Header) http.Header { diff --git a/core/http/middleware/trace_redact_test.go b/core/http/middleware/trace_redact_test.go index 6f1a83586f72..787a30678128 100644 --- a/core/http/middleware/trace_redact_test.go +++ b/core/http/middleware/trace_redact_test.go @@ -13,6 +13,12 @@ import ( // any heap inspection. This pins the redaction contract so a future refactor // of TraceMiddleware can't silently regress it. var _ = Describe("redactSensitiveHeaders", func() { + It("redacts WebSocket ticket subprotocols", func() { + h := http.Header{} + h.Set("Sec-WebSocket-Protocol", "localai.motion.v2, localai.ticket.secret") + Expect(redactSensitiveHeaders(h).Get("Sec-WebSocket-Protocol")).To(Equal("[redacted]")) + }) + It("redacts Authorization", func() { h := http.Header{} h.Set("Authorization", "Bearer sk-secret-1234567890") diff --git a/core/http/react-ui/e2e/discover-height.spec.js b/core/http/react-ui/e2e/discover-height.spec.js index a03feb0c9752..8a0fb47fc331 100644 --- a/core/http/react-ui/e2e/discover-height.spec.js +++ b/core/http/react-ui/e2e/discover-height.spec.js @@ -23,12 +23,20 @@ test.describe('Models Explore - the view scrolls, not the page', () => { test.beforeEach(async ({ page }) => { await page.route('**/api/models*', (route) => route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK) })) + await page.route('**/api/resources', (route) => + route.fulfill({ + contentType: 'application/json', + body: JSON.stringify({ aggregate: { total_memory: 12 * 1024 ** 3, gpu_count: 1 } }), + })) }) test('a long detail scrolls the pane and leaves the page height alone', async ({ page }) => { await page.setViewportSize({ width: 1400, height: 900 }) await page.goto('/app/models') await expect(page.locator('[data-testid="discover-rail-item"]').first()).toBeVisible({ timeout: 10_000 }) + // Resources load independently of the gallery; the fits toggle takes a + // row from the rail, so include it before measuring selection's effect. + await expect(page.getByText('Fits in GPU', { exact: true })).toBeVisible() const pageHeight = () => page.evaluate(() => document.documentElement.scrollHeight) const railHeight = () => page.evaluate( diff --git a/core/http/react-ui/public/locales/en/importModel.json b/core/http/react-ui/public/locales/en/importModel.json index 0a934f94f747..9795060da662 100644 --- a/core/http/react-ui/public/locales/en/importModel.json +++ b/core/http/react-ui/public/locales/en/importModel.json @@ -81,6 +81,7 @@ "video": "Video generation", "3d": "3D mesh generation", "3d_animation": "3D animation", + "motion": "Motion capture", "embeddings": "Embeddings", "reranker": "Rerankers", "detection": "Object detection", diff --git a/core/http/react-ui/public/locales/en/models.json b/core/http/react-ui/public/locales/en/models.json index b4f4516c2dba..8559e788962c 100644 --- a/core/http/react-ui/public/locales/en/models.json +++ b/core/http/react-ui/public/locales/en/models.json @@ -3,7 +3,8 @@ "title": "Models", "navLabel": "Model lifecycle", "views": { "explore": "Explore", "installed": "Installed" }, - "filters": { "all": "All", "running": "Running", "idle": "Idle", "disabled": "Disabled", "pinned": "Pinned", "distributed": "Distributed" }, + "filters": { + "motion": "Motion Capture", "all": "All", "running": "Running", "idle": "Idle", "disabled": "Disabled", "pinned": "Pinned", "distributed": "Distributed" }, "installed": { "searchPlaceholder": "Search installed models", "count": "{{shown}} of {{total}}", @@ -98,6 +99,7 @@ "delete": "Delete" }, "filters": { + "motion": "Motion Capture", "all": "All", "llm": "Chat", "image": "Image", diff --git a/core/http/react-ui/src/pages/ImportModel.jsx b/core/http/react-ui/src/pages/ImportModel.jsx index 259bde7b5336..3878db16e481 100644 --- a/core/http/react-ui/src/pages/ImportModel.jsx +++ b/core/http/react-ui/src/pages/ImportModel.jsx @@ -25,7 +25,7 @@ const BACKENDS_FALLBACK_EMPTY = [] // Modality keys used as i18n keys under "modality.*" namespace; resolved // at render time inside `buildBackendOptions`. -const MODALITY_KEYS = ['text', 'asr', 'tts', 'image', 'video', '3d', '3d_animation', 'embeddings', 'reranker', 'detection', 'vad'] +const MODALITY_KEYS = ['text', 'asr', 'tts', 'image', 'video', '3d', '3d_animation', 'motion', 'embeddings', 'reranker', 'detection', 'vad'] // buildBackendOptions groups known backends by modality and tags // auto_detect=false entries with a muted "manual pick" badge so users diff --git a/core/http/react-ui/src/pages/Models.jsx b/core/http/react-ui/src/pages/Models.jsx index 6479011c12ae..0c49ab1bce63 100644 --- a/core/http/react-ui/src/pages/Models.jsx +++ b/core/http/react-ui/src/pages/Models.jsx @@ -95,6 +95,7 @@ const FILTERS = [ { key: 'video', labelKey: 'filters.video', icon: 'video' }, { key: '3d', labelKey: 'filters.threed', icon: 'cube' }, { key: '3d_animation', labelKey: 'filters.threedAnimation', icon: 'walk' }, + { key: 'motion', labelKey: 'filters.motion', icon: 'walk' }, { key: 'multimodal', labelKey: 'filters.multimodal', icon: 'shapes' }, { key: 'vision', labelKey: 'filters.vision', icon: 'eye' }, { key: 'tts', labelKey: 'filters.tts', icon: 'mic' }, diff --git a/core/http/react-ui/src/utils/capabilities.js b/core/http/react-ui/src/utils/capabilities.js index 722c85842062..f0cd3a3883f6 100644 --- a/core/http/react-ui/src/utils/capabilities.js +++ b/core/http/react-ui/src/utils/capabilities.js @@ -31,3 +31,5 @@ export const CAP_REALTIME_AUDIO = 'FLAG_REALTIME_AUDIO' export const CAP_SCORE = 'FLAG_SCORE' export const CAP_DECISIONS = 'FLAG_DECISIONS' export const CAP_TOKEN_CLASSIFY = 'FLAG_TOKEN_CLASSIFY' + +export const CAP_MOTION = 'FLAG_MOTION' diff --git a/core/http/route_coverage_test.go b/core/http/route_coverage_test.go index 65dfe37f497a..61c2db65a9de 100644 --- a/core/http/route_coverage_test.go +++ b/core/http/route_coverage_test.go @@ -73,6 +73,19 @@ var _ = Describe("Route auth coverage", func() { Expect(os.RemoveAll(tmpdir)).To(Succeed()) }) + It("exposes only session controls and duplex streaming for motion", func() { + var motionRoutes []string + for _, route := range app.Routes() { + if strings.HasPrefix(route.Path, "/api/motion/") { + motionRoutes = append(motionRoutes, route.Method+" "+route.Path) + } + } + Expect(motionRoutes).To(ConsistOf( + "POST /api/motion/sessions", "GET /api/motion/sessions/:id", "DELETE /api/motion/sessions/:id", + "POST /api/motion/sessions/:id/tickets", "GET /api/motion/sessions/:id/poses", + )) + }) + It("enforces the anonymous-access decision for every registered route", func() { type routePattern struct { method string diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index 4a78e30e4769..8450ea582564 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -41,6 +41,13 @@ func RegisterLocalAIRoutes(router *echo.Echo, c.URLs = []string{"doc.json"} })) + motion := localai.NewMotionEndpoints(app) + router.POST("/api/motion/sessions", motion.Create) + router.GET("/api/motion/sessions/:id", motion.Get) + router.DELETE("/api/motion/sessions/:id", motion.Delete) + router.GET("/api/motion/sessions/:id/poses", motion.Poses) + router.POST("/api/motion/sessions/:id/tickets", motion.Ticket) + // LocalAI API endpoints if !appConfig.DisableGalleryEndpoint { // Import model page diff --git a/core/http/routes/ui_api.go b/core/http/routes/ui_api.go index 08ced3bc0e75..e9a51af19132 100644 --- a/core/http/routes/ui_api.go +++ b/core/http/routes/ui_api.go @@ -45,6 +45,7 @@ const ( // usecaseFilters maps UI filter keys to ModelConfigUsecase flags for // capability-based gallery filtering. var usecaseFilters = map[string]config.ModelConfigUsecase{ + config.UsecaseMotion: config.FLAG_MOTION, config.UsecaseChat: config.FLAG_CHAT, config.UsecaseImage: config.FLAG_IMAGE, config.UsecaseVideo: config.FLAG_VIDEO, diff --git a/core/services/nodes/health_mock_test.go b/core/services/nodes/health_mock_test.go index fc00e6bb1b53..95ea866d7bad 100644 --- a/core/services/nodes/health_mock_test.go +++ b/core/services/nodes/health_mock_test.go @@ -257,6 +257,10 @@ func (c *fakeBackendClient) AudioToAudioStream(_ context.Context, _ ...ggrpc.Cal func (c *fakeBackendClient) AudioTranscriptionLive(_ context.Context, _ ...ggrpc.CallOption) (grpc.AudioTranscriptionLiveClient, error) { return nil, nil } +func (c *fakeBackendClient) MotionStream(_ context.Context, _ ...ggrpc.CallOption) (grpc.MotionStreamClient, error) { + return nil, fmt.Errorf("motion streaming not implemented by health fake") +} + func (c *fakeBackendClient) Forward(_ context.Context, _ ...ggrpc.CallOption) (grpc.ForwardClient, error) { return nil, nil } diff --git a/core/services/nodes/inflight_test.go b/core/services/nodes/inflight_test.go index 77b7314e03b5..1c2f6df209bd 100644 --- a/core/services/nodes/inflight_test.go +++ b/core/services/nodes/inflight_test.go @@ -218,6 +218,10 @@ func (f *fakeGRPCBackend) AudioTranscriptionLive(_ context.Context, _ ...ggrpc.C return nil, nil } +func (f *fakeGRPCBackend) MotionStream(_ context.Context, _ ...ggrpc.CallOption) (grpc.MotionStreamClient, error) { + return nil, fmt.Errorf("motion streaming not implemented by in-flight fake") +} + func (f *fakeGRPCBackend) Forward(_ context.Context, _ ...ggrpc.CallOption) (grpc.ForwardClient, error) { return nil, nil } diff --git a/core/services/worker/free_timeout_test.go b/core/services/worker/free_timeout_test.go index 4f1b6346e749..70f470ea5362 100644 --- a/core/services/worker/free_timeout_test.go +++ b/core/services/worker/free_timeout_test.go @@ -4,6 +4,7 @@ import ( "context" "net" "os" + "os/exec" "strconv" "syscall" @@ -82,9 +83,11 @@ var _ = Describe("Stopping a backend whose Free never returns", func() { // actually dead afterwards, not merely that Stop() returned. It // outlives every timeout below, so if it is gone at the end it is // because the supervisor signalled it. + sleepPath, err := exec.LookPath("sleep") + Expect(err).NotTo(HaveOccurred()) proc = process.New( process.WithTemporaryStateDir(), - process.WithName("/bin/sleep"), + process.WithName(sleepPath), process.WithArgs("300"), ) Expect(proc.Run()).To(Succeed()) diff --git a/core/services/worker/model_stop_test.go b/core/services/worker/model_stop_test.go index 65ddaf20de14..7deb55ef6f27 100644 --- a/core/services/worker/model_stop_test.go +++ b/core/services/worker/model_stop_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "net" + "os/exec" "sync/atomic" process "github.com/mudler/go-processmanager" @@ -37,7 +38,9 @@ func startModelStopBackend(backend *modelStopBackend) (string, int, func()) { } func startModelStopProcess() *process.Process { - proc := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sleep"), process.WithArgs("300")) + sleepPath, err := exec.LookPath("sleep") + Expect(err).NotTo(HaveOccurred()) + proc := process.New(process.WithTemporaryStateDir(), process.WithName(sleepPath), process.WithArgs("300")) Expect(proc.Run()).To(Succeed()) return proc } diff --git a/core/trace/backend_trace.go b/core/trace/backend_trace.go index e14370c455f8..5cf6299bc70a 100644 --- a/core/trace/backend_trace.go +++ b/core/trace/backend_trace.go @@ -31,6 +31,7 @@ const ( BackendTrace3DGeneration BackendTraceType = "3d_generation" BackendTrace3DRemesh BackendTraceType = "3d_remesh" BackendTrace3DAnimation BackendTraceType = "3d_animation" + BackendTraceMotion BackendTraceType = "motion" BackendTraceTTS BackendTraceType = "tts" BackendTraceSoundGeneration BackendTraceType = "sound_generation" BackendTraceRerank BackendTraceType = "rerank" diff --git a/docs/content/features/backends.md b/docs/content/features/backends.md index a9afd0ce14cf..db117cb5be8d 100644 --- a/docs/content/features/backends.md +++ b/docs/content/features/backends.md @@ -349,3 +349,7 @@ has actually gone away; a client that waits receives the full context worth of tokens. Set `max_tokens` on the model config, and keep `repeat_penalty` above `1` so a repetition loop terminates on its own. {{% /notice %}} + +### Motion capture + +- [gem-x.cpp](https://github.com/localai-org/gem-x.cpp): live SOMA-77 and SMPL-24 human poses on CPU/Vulkan. See [Motion Capture](/features/motion/). diff --git a/docs/content/features/motion.md b/docs/content/features/motion.md new file mode 100644 index 000000000000..cb416a6b17ea --- /dev/null +++ b/docs/content/features/motion.md @@ -0,0 +1,362 @@ ++++ +title = "Motion Capture" +url = "/features/motion/" +weight = 21 ++++ + +LocalAI's `gemxcpp` backend converts timestamped RGB frames into human poses +using [gem-x.cpp](https://github.com/localai-org/gem-x.cpp). It runs resident +GEM-X, ViTPose and YOLOX models on CPU or Vulkan. Live inference does not need +SAM3D Body. This initial integration is API-only: no Motion UI, offline clip +export, body mesh or SONIC controller is included. + +## Install and configure + +Import `https://huggingface.co/LocalAI-io/GEM-X-GGUF` through the model importer. +It selects `gemxcpp` and downloads three checksummed, revision-pinned assets. +These are custom F32 GGUF architectures, not language models. Weight licenses +are described in the source model card. Allow memory for all three models and +execution buffers, beyond their roughly 4 GB combined download size. + +Equivalent configuration with existing local files: + +```yaml +name: gem-x +backend: gemxcpp +known_usecases: [motion] +parameters: + model: gem-x/gem-x-contact-f32.gguf +threads: 4 +options: + - vitpose:gem-x/vitpose-f32.gguf + - yolox:gem-x/yolox-f32.gguf + - window:30 + - detector_interval:1 + - selection:continuity + - precision:strict + - max_gap_us:2000000 +``` + +Additional options: `device:cpu` or `device:vulkan`, and `device_index:0`. +Without `device`, the package selects CPU or Vulkan according to its build. +Darwin uses CPU; Metal inference is not supported. Explicit unavailable devices +fail rather than silently falling back. Threads are capped at eight. + +`window` accepts 2–120 observations; `detector_interval` accepts 1–30 processed +frames. Larger detector intervals reuse the previous crop and may miss movement +or person loss until the next detection. `selection:continuity` conservatively +associates one unique box; loss/ambiguity clears history. `selection:parity` +uses upstream's largest detection/full-image fallback and makes no identity +continuity guarantee. Neither mode is person re-identification. + +`precision:strict` is the default. The launcher fixes GGML Vulkan precision +flags before initialization. `precision:backend_default` opts out of strict +validation, without promising a faster path. The same F32 model files are used. + +## Sessions + +All endpoints use standard LocalAI authentication, the **Motion Capture** feature +permission and model access controls. Sessions belong to their creating user. +When authentication is disabled, sessions share the unauthenticated principal. +Legacy admin keys have the existing shared legacy-admin identity. + +```sh +curl http://localhost:8080/api/motion/sessions \ + -H 'Content-Type: application/json' \ + -d '{"model":"gem-x","profile":"smpl24","time_origin":"camera session start"}' +``` + +The response is HTTP 201 with `id`, `model`, `profile`, `time_origin` and counters. +`profile` defaults to `soma77`; `time_origin` is required (1–128 characters). +A session owns its model's native temporal state. Only one session per model +is admitted, with up to 16 sessions per frontend and one duplex connection per +session. Additional model configurations consume independent model memory. + +| Endpoint | Purpose | +| --- | --- | +| `POST /api/motion/sessions` | Create and load a session | +| `GET /api/motion/sessions/{id}` | Read frame/pose/drop counters | +| `DELETE /api/motion/sessions/{id}` | Cancel and release a session | +| `GET /api/motion/sessions/{id}/poses` | Exchange frames and poses over a duplex WebSocket | +| `POST /api/motion/sessions/{id}/tickets` | Issue a single-use browser WebSocket ticket | + +Sessions expire after two minutes without frame activity. Inference also +has a two-minute timeout; initial model loading has a five-minute stream wait +limit, subject to the loader's own cancellation behavior. Native inference is +synchronous and cannot be interrupted mid-call. Closing stops new work and +discards late results; native cleanup completes after the current call returns. + +Session state lives in one frontend process. Use sticky routing for all session +requests in multi-frontend deployments. Backend streaming uses the existing +model loader and keeps its selected gRPC connection for the session. There is +no session migration or recovery after frontend/backend restart. + +## Duplex frame transport (v2) + +The browser negotiates `localai.motion.v2` on the existing ticket-authenticated +`/poses` socket. It sends uncompressed binary protobuf `Input.frame` and receives +binary `Output` messages. HTTP remains for discovery, session creation, tickets +and deletion. There is no HTTP frame fallback. The protobuf package remains v1; +the v2 name versions the duplex transport and flow protocol. + +Definition precedes a JSON flow hello. Each flow message has `type: "flow"`, +`version: 2`, a consecutive integer `id`, `release` containing decimal sequence +strings, `window_frames`, `max_pending_bytes`, `max_frame_bytes`, +`processing_ewma_ms`, `recommended_fps`, `queued_frames`, `queued_bytes`, +`oldest_age_ms`, `pressure` and cumulative `dropped_frames`. A completed frame +also has `completed_sequence` and `completed_kind` (`pose` or `event`). Its binary +output is written before its release; discarded inputs have release only. +While connected, every completed or dropped capture is retired once, and pose identity retains source time. + +The initial window is two frames, with zero recommended FPS meaning camera-paced +subject to credits. After a service sample the window is eight. Service EWMA uses +alpha 0.2 and a 1ms floor; normal offered rate is 1.05/service, pressure rate is +0.9/service. Clients should gate capture before canvas/RGB allocation on credits, +outstanding bytes, pacing and prospective WebSocket upload backlog. A recommended +upload backlog limit is two frame sizes or 2MiB, whichever is smaller. Prefer fresh +camera callbacks over a fixed capture timer or a local RGB FIFO. Camera/upload +speed can limit the achieved rate. These are recommendations for external clients; +LocalAI does not include a browser capture client or enforce client-side pacing. + +LocalAI uses one reader, one inference worker and one writer. Pending input is +bounded to eight frames/8MiB with a 2MiB encoded frame limit. At six frames or +6MiB, thinning retains half the pending frames nearest evenly spaced source-time +targets across oldest/newest (earlier ties; one survivor means newest). Active +inference is never thinned. Inputs older than 250ms since server admission expire +before enqueue/dispatch. Pressure clears once pending depth is at most two. +The 8MiB pending limit is not a total process-memory claim: encoded/decoded +active and reader payloads add up to approximately another 6MiB, excluding runtime, +protobuf and transport overhead. The writer queue holds at most eight messages; +writes time out at five seconds. External clients should also detect stalled +feedback; a five-second timeout is recommended. + +A connection exclusively owns ingress: a second connection receives HTTP 409. +Only `localai.motion.v2` is accepted; missing or read-only subprotocols receive HTTP 400. The socket accepts +frames only; reset/camera/model changes require a new session. Disconnect cancels +transport, discards pending input and joins workers before releasing ownership; +this does not preempt an executing native GPU kernel. Only processed frames enter +existing `motion_frames_v1` accounting; dropped frames are not billed. Session +summaries expose upload accepted/dropped counts; ingress metrics distinguish +accepted, thinned and expired outcomes. Upload drop counts exclude pending/active +inputs disposed by disconnect; disconnect clears all browser credits. No per-frame content or credentials are logged. + +## Frame schema and authentication + +The authoritative schema is `pkg/motion/proto/motion.proto` in the LocalAI +repository, package `localai.motion.v1`. Generate client bindings from that file; +clients do not need LocalAI's internal backend protocol. Numeric arrays are +packed float32, joint-major then component. JavaScript clients must retain +64-bit sequence/timestamp fields as BigInt or their Protobuf library's integer +representation, not lossy JavaScript numbers for arbitrarily large values. + +1. Create the session through HTTP JSON. +2. Connect to `/api/motion/sessions/{id}/poses` over `ws://` or `wss://`, + requesting subprotocol `localai.motion.v2`. Native clients can use Bearer + authentication; browsers use their authenticated same-origin session cookie. + Successfully validated `Authorization`, `x-api-key` or `xi-api-key` header + credentials bypass the WebSocket origin check. Invalid or unvalidated headers + do not; cookie-only authentication still follows the origin policy. Standard + browser WebSocket APIs cannot set these headers, so browser clients using + cookies need an explicit CORS policy: set `LOCALAI_CORS=true` + and `LOCALAI_CORS_ALLOW_ORIGINS=https://consumer.example,http://localhost:3000`. + This explicit origin policy also permits cross-site session creation, tickets + and deletion through CSRF protection. Authentication and permissions + still apply. Allowed origins use the HTTP CORS matching rules, including explicitly + configured wildcards. Without an explicit policy or validated header credentials, + sockets remain same-origin even though HTTP uses a permissive default. Native clients without an Origin + header are supported. Origin approval does not bypass authentication or session + ownership; browser cookie availability still depends on cookie and browser + policies. Credentials are not URL parameters. +3. Read the first binary message as `Output.definition`. It declares joint order, + parents, rest transforms, coordinate conventions, profile and available channels. + The definition is sent once per session connection. A disconnect ends the session; + create a new session before reconnecting. +4. Send serialized `Input.frame` messages as binary WebSocket messages, following + the flow credits described above. A frame contains tightly packed RGB8 + bytes, width/height, sequence, source microseconds and optional subject box/ID. + Subject boxes use inclusive source pixel XYXY coordinates, bounded by + `[0, width-1]` and `[0, height-1]`. Frames must be at least 8 pixels per side, at most 32766 per side and at most + 16 million pixels. RGB byte length must equal width × height × 3. The stricter + 2MiB encoded-message limit also applies to every upload. +5. Read `Output.pose` or `Output.event`, followed by the JSON completion release. + The first usable observation produces `warmup`; at least two are needed for a pose. + +Invalid frames, replayed sequences and non-increasing timestamps close the +connection without submitting that input to the backend. Source timestamps are +nonnegative microseconds relative to the declared origin and must strictly +increase, as must sequence numbers. Gaps in sequence numbers are allowed. +Create a new session before seeking backward, explicitly resetting or changing +cameras. Long source gaps, changed dimensions or a changed caller-selected +subject reset native context; poses/events carry the new epoch and reset flag. +Intrinsics follow upstream's image-centre/max-dimension approximation; +calibrated cameras are not supported by this ABI. + +The model uses accepted-frame indices, not time-aware irregular sampling. Source +timestamps are preserved, not replaced by receive times or an invented FPS. +Consumers own resampling and clock alignment. Local processing durations do not +measure capture-to-consumer latency. + +Pose and lifecycle outputs are delivered in order through the bounded writer +queue. A stalled reader eventually causes a write timeout and session closure; +outputs are not replaced by newer poses. WebSocket uses TCP, so already-buffered +bytes cannot be retracted. Reconnecting requires a new session; there is no +history replay. + +## Browser WebSocket tickets + +Browsers cannot set `Authorization` on a native `WebSocket` constructor. After +creating a motion session, exchange your normal credentials for a ticket: + +```js +const response = await fetch(`${apiBase}/api/motion/sessions/${sessionId}/tickets`, { + method: "POST", + headers: { + "Authorization": `Bearer ${apiKey}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ origin: window.location.origin }), +}); +if (!response.ok) throw new Error(`Ticket request failed: ${response.status}`); +const { ticket, expires_at } = await response.json(); +const wsBase = apiBase.replace(/^http/, "ws"); +const socket = new WebSocket(`${wsBase}/api/motion/sessions/${sessionId}/poses`, [ + "localai.motion.v2", + `localai.ticket.${ticket}`, +]); +socket.binaryType = "arraybuffer"; +``` + +Cookie-authenticated clients can use `credentials: "include"` instead of the +Authorization header, subject to browser cookie policy. The ticket endpoint uses +normal authentication, Motion feature permission, model access and session +ownership checks. Its JSON response is HTTP 201 with `ticket` and `expires_at`, +and `Cache-Control: no-store`. When auth is disabled, tickets retain the existing +anonymous access policy; they do not create a user or grant additional access. + +Tickets expire after **30 seconds**, are **single use**, and are bound to the +exact browser origin and session's pose WebSocket path. If the HTTP request has +an Origin header, it must match the requested origin. A reconnect requires a new +session and ticket. Redemption consumes the ticket before authentication/upgrade; even a +failed matching upgrade requires a fresh ticket. The original credentials are +revalidated, so revoked keys or sessions cannot redeem outstanding tickets. +Feature, model and ownership permissions are checked again. Successfully redeemed +tickets authorize the bound origin without requiring browser cookie delivery. + +The server selects only `localai.motion.v2`; it never echoes the ticket protocol. +Tickets belong in the protocol list, not in the URL or query string. LocalAI +strips them before upgrade and redacts the protocol header from API traces. +Configure external proxies and request loggers to redact `Sec-WebSocket-Protocol` +as well. Ticket responses must not be cached or logged. + +There are at most 32 outstanding tickets per user (shared for anonymous/legacy +identities) and 4096 per frontend; issuance returns 429 at capacity. Closing a +session revokes its outstanding tickets. The in-memory store is shared by ticket +issuance and upgrade on one frontend; both requests need sticky routing. Restart +invalidates all outstanding tickets. Ticket expiry limits the handshake window, +not the lifetime of an established connection. Ticket requests do not add frame +usage. This first integration supports motion; other WebSocket endpoints do not +yet issue tickets. + +## Output profiles + +**soma77:** 77 joints, metres, right-handed native SOMA Y-up. Positions have pelvis +translation removed only; root rotation is retained. Local rotations are XYZW; +local translations vary with the emitted frame's identity/scale. Root axis-angle +is in radians. Root translation is zero: this is not a continuous world trajectory. + +**smpl24:** upstream-validated SOMA-to-SMPL mapping; 24 joint positions with pelvis +translation and anchor rotation removed. The separate anchor is a unit WXYZ +quaternion in the upstream Z-up/base-rotation convention. It reconstructs the +mapped, gravity-aligned points from root-local positions. No SMPL joint rotations, +mesh parameters, physical-base alignment or robot wrist commands are fabricated. + +Definitions distinguish neutral reference transforms from per-frame shape. +Quaternions use upstream's canonical nonnegative-w policy, not temporal sign +smoothing; consumers interpolating rotations should handle equivalent signs. + +Events include `warmup`, `lost`, `ambiguous` and `reset`. Frame flags identify +reset (1), reused crop (2), full-image fallback (4), detector execution (8) and +caller box (16). Track epochs mark acquisitions; automatic parity mode uses +track epoch zero without asserting identity continuity. No 3D confidence is +advertised. Consumers must detect stale streams and choose their own hold, +interpolation or stop behavior. + +An external SONIC consumer can request SMPL-24, buffer timed poses and construct +its policy-specific reference windows. LocalAI does not execute that policy, +align references to physical robot state, or run physics. + +## Observability + +With metrics enabled, the existing `/metrics` exporter includes: + +- `localai_motion_sessions`: active sessions. +- `localai_motion_frames`: processed frame outcomes. +- `localai_motion_uploads`: accepted, thinned and expired input frames. +- `localai_motion_resets`: explicit or native context resets. +- `localai_motion_frame_duration_seconds`: backend frame round-trip duration. +- `localai_motion_stage_duration_seconds`: detector, ViTPose, GEM and skeleton + construction stage durations from the native pipeline. + +Prometheus may append standard counter/unit suffixes. Labels use model/backend, +finite outcome names and fixed stage names, never session IDs or person IDs. +Frame/pose/drop totals are also available through session status. + +Frame submissions also use the shared usage accounting pipeline, like animation +requests. Each successfully processed frame records `input_units: 1`; a pose +records `output_units: 1`, while warmup, lost and ambiguous outcomes record zero +output units. Records retain `accounting_rule: "motion_frames_v1"` in their usage +metadata. The existing usage dashboard, storage and billing metrics expose these +units through their legacy prompt/completion/total token fields; they represent +frames, not text tokens. Configure any model pricing in those units. + +Accounting happens once per successfully processed WebSocket frame, before +socket delivery. Rejected or dropped inputs and failed inference do not record +usage. Session creation/status/deletion, tickets and WebSocket upgrades do not +record usage either. Shared statistics can be disabled through the existing +statistics configuration. + +With backend tracing enabled, a `motion` trace spans the session and records its +profile, duration, termination reason and counts. Model-load failures use the +existing model-load traces. Logs record lifecycle changes. Images, pose arrays, +credentials and private model paths are not included in motion session traces. + +### Webcam pose overlays + +GEM-X pose messages optionally include `image_positions` (Protobuf field 13): +packed source-image pixel XY pairs in `Definition.joint_names` order (48 values for +SMPL24, 154 for SOMA77). The definition advertises the channel and convention +`source-pixel xy, joint_names order, top-left origin, unmirrored`. Coordinates +reference the exact submitted RGB Frame identified by the pose sequence/timestamp; +use that frame's width/height, then scale/letterbox with the displayed image. They +are neither normalized crop coordinates nor confidence scores. No image mirroring +is applied. Pixels can fall outside the source image; clients should clip drawing. + +The backend projects GEM-X's emitted camera-space skeleton plus its camera +translation using the native live API's actual intrinsics: focal length +`max(width,height)` and image-centre principal point. SMPL uses the mapping from +that native definition, not a hard-coded substitute skeleton. If a joint has +nonpositive depth or projection is nonfinite the optional channel is empty; root +local control data remains separate. Older servers may omit the channel. Do not +invent an image-aligned overlay from root-local positions, anchor or bounding box. + +### Optional live root displacement + +The GEM-X SMPL stream advertises `root_displacement` only when its native library +supports channel 14. It contains three metres-per-interval components in the +**end pose's anchor-local basis**, not velocity or absolute world position. Rotate +by the end pose's WXYZ anchor for gravity-aligned Z-up displacement. The closed +interval is `[displacement_start_time_us, source_time_us]`, between consecutive +frames actually processed by GEM-X (including its warmup observation). + +This uses the last two decoded world translations inside one rolling window; +it does not subtract absolute translations from separate windows. The current +frame's unobserved next-interval prediction is not emitted. `root_translation` +remains zero for compatibility. Older native libraries omit the new channel. + +Consumers must establish a baseline after an epoch/track change, reset, missing +interval or reconnect; they must not add a delta across those boundaries. A +dropped input changes the model's accepted-frame spacing: these estimates are +not time-normalized or contact-refined, and integration can drift. Robot/world +placement is a separate consumer-owned transform and need not reset on tracking +loss. Model-backed checks cover interval boundaries and coordinate-basis +consistency; they do not establish trajectory accuracy against ground truth. diff --git a/docs/content/features/runtime-settings.md b/docs/content/features/runtime-settings.md index 842552ec5423..4aa34fab4c89 100644 --- a/docs/content/features/runtime-settings.md +++ b/docs/content/features/runtime-settings.md @@ -91,8 +91,8 @@ You can configure these settings via the web UI or through environment variables ### API Security - **CORS**: Enable Cross-Origin Resource Sharing -- **CORS Allow Origins**: Comma-separated list of allowed CORS origins -- **CSRF**: Enable CSRF protection middleware. Cross-site browser requests are exempt only when the server successfully authenticates a credential from `Authorization`, `x-api-key`, or `xi-api-key`. Supplying an arbitrary header when authentication is disabled does not bypass protection; cookie authentication also does not grant this exemption. Requests without `Sec-Fetch-Site` (such as CLI clients) remain allowed by this middleware. +- **CORS Allow Origins**: Comma-separated list of allowed CORS origins. When CORS is explicitly enabled with a nonempty allowlist, matching origins are also trusted by CSRF protection for cross-site state-changing requests. Explicit wildcard patterns apply to both policies. The default permissive HTTP CORS policy does not bypass CSRF; authentication and permissions still apply. +- **CSRF**: Enable CSRF protection middleware. Outside explicitly trusted CORS origins, cross-site browser requests are exempt only when the server successfully authenticates a credential from `Authorization`, `x-api-key`, or `xi-api-key`. Supplying an arbitrary header when authentication is disabled does not bypass protection; cookie authentication also does not grant this exemption. Requests without `Sec-Fetch-Site` (such as CLI clients) remain allowed by this middleware. - **API Keys**: Manage API keys for authentication (one per line or comma-separated) For multi-user authentication with roles, OAuth, and usage tracking, see [Authentication & Authorization]({{%relref "features/authentication" %}}). diff --git a/gallery/index.yaml b/gallery/index.yaml index a68ad95e59c0..0cd1837b5f75 100644 --- a/gallery/index.yaml +++ b/gallery/index.yaml @@ -69809,3 +69809,36 @@ - filename: llama-cpp/models/diarizationlm-gemma-4-e4b-v1/DiarizationLM-Gemma-4-E4B-v1-q4_0.gguf uri: https://huggingface.co/google/DiarizationLM-Gemma-4-E4B-v1/resolve/91b7c06adae025090dc7fcec3afcc1cb145c7302/DiarizationLM-Gemma-4-E4B-v1-q4_0.gguf sha256: 297238b1dbfe8bb061f5e4cc79be0e2d607a4539d7bc6dc00d75cc0428f6a264 + +- name: gem-x + url: "github:mudler/LocalAI/gallery/virtual.yaml@master" + urls: + - https://huggingface.co/LocalAI-io/GEM-X-GGUF + description: | + Live human motion capture with GEM-X on CPU/Vulkan. Streams SOMA-77 or + SMPL-24 poses through /api/motion; no body mesh or SAM3D dependency. + tags: + - motion + - pose-estimation + - gguf + overrides: + backend: gemxcpp + known_usecases: [motion] + parameters: + model: gem-x/gem-x-contact-f32.gguf + options: + - vitpose:gem-x/vitpose-f32.gguf + - yolox:gem-x/yolox-f32.gguf + - selection:continuity + - window:30 + - detector_interval:1 + files: + - filename: gem-x/gem-x-contact-f32.gguf + sha256: 175857b8453b65d028f3703d813c3186b31b66ea65fbdb94e64308c39e083f78 + uri: https://huggingface.co/LocalAI-io/GEM-X-GGUF/resolve/b36180fd4c0d7c6a7fb2269348f632d3ed279868/gem-x-contact-f32.gguf + - filename: gem-x/vitpose-f32.gguf + sha256: 272c75d4c3a6a740f1eb3f3222832de206150fa035c134af17e66ca9f8386d11 + uri: https://huggingface.co/LocalAI-io/GEM-X-GGUF/resolve/b36180fd4c0d7c6a7fb2269348f632d3ed279868/vitpose-f32.gguf + - filename: gem-x/yolox-f32.gguf + sha256: 2be3d28e0dd8a171f4ad980b3dfbddaa1ecb88e478c612bc11379913908e7b78 + uri: https://huggingface.co/LocalAI-io/GEM-X-GGUF/resolve/b36180fd4c0d7c6a7fb2269348f632d3ed279868/yolox-f32.gguf diff --git a/pkg/grpc/backend.go b/pkg/grpc/backend.go index ac5b95f28efe..86f4b862860b 100644 --- a/pkg/grpc/backend.go +++ b/pkg/grpc/backend.go @@ -124,6 +124,7 @@ type ControlBackend interface { AudioTransformStream(ctx context.Context, opts ...grpc.CallOption) (AudioTransformStreamClient, error) AudioToAudioStream(ctx context.Context, opts ...grpc.CallOption) (AudioToAudioStreamClient, error) AudioTranscriptionLive(ctx context.Context, opts ...grpc.CallOption) (AudioTranscriptionLiveClient, error) + MotionStream(ctx context.Context, opts ...grpc.CallOption) (MotionStreamClient, error) // Forward proxies a raw HTTP request to an upstream provider for // passthrough-mode cloud-proxy backends. Caller streams a single diff --git a/pkg/grpc/base/base.go b/pkg/grpc/base/base.go index ace2a53fa300..2c8f52329008 100644 --- a/pkg/grpc/base/base.go +++ b/pkg/grpc/base/base.go @@ -251,3 +251,7 @@ func memoryUsage() *pb.MemoryUsageData { func (llm *Base) Free() error { return nil } + +func (*Base) MotionStream(context.Context, func() (*pb.MotionRequest, error), func(*pb.MotionResponse) error) error { + return fmt.Errorf("motion streaming is not supported") +} diff --git a/pkg/grpc/channel_stream.go b/pkg/grpc/channel_stream.go new file mode 100644 index 000000000000..f8d9c3e9de0e --- /dev/null +++ b/pkg/grpc/channel_stream.go @@ -0,0 +1,129 @@ +// SPDX-License-Identifier: MIT +package grpc + +import ( + "context" + "io" + "sync" + + "google.golang.org/grpc/metadata" +) + +// channelStream carries in-process bidirectional RPCs. Publishing the terminal +// error before closing responses lets callers drain the response tail reliably. +type channelStream[Request, Response any] struct { + ctx context.Context + requests chan Request + responses chan Response + done chan struct{} + err error +} + +type channelStreamServer[Request, Response any] struct { + *channelStream[Request, Response] +} + +type channelStreamClient[Request, Response any] struct { + *channelStream[Request, Response] + sendMu sync.Mutex + closed bool + cleanup *streamCleanup +} + +func newChannelStream[Request, Response any](ctx context.Context, capacity int, serve func(*channelStreamServer[Request, Response]) error) *channelStreamClient[Request, Response] { + stream := &channelStream[Request, Response]{ctx: ctx, requests: make(chan Request, capacity), responses: make(chan Response, capacity), done: make(chan struct{})} + client := &channelStreamClient[Request, Response]{channelStream: stream, cleanup: newStreamCleanup(ctx, nil)} + go func() { + stream.err = serve(&channelStreamServer[Request, Response]{stream}) + close(stream.done) + close(stream.responses) + }() + return client +} + +func (s *channelStream[Request, Response]) Context() context.Context { return s.ctx } +func (s *channelStreamServer[Request, Response]) Send(response Response) error { + select { + case s.responses <- response: + return nil + case <-s.ctx.Done(): + return s.ctx.Err() + } +} +func (s *channelStreamServer[Request, Response]) Recv() (Request, error) { + select { + case request, ok := <-s.requests: + if ok { + return request, nil + } + case <-s.ctx.Done(): + var zero Request + return zero, s.ctx.Err() + } + var zero Request + return zero, io.EOF +} +func (s *channelStreamServer[Request, Response]) SetHeader(metadata.MD) error { return nil } +func (s *channelStreamServer[Request, Response]) SendHeader(metadata.MD) error { return nil } +func (s *channelStreamServer[Request, Response]) SetTrailer(metadata.MD) {} +func (s *channelStreamServer[Request, Response]) SendMsg(message any) error { + if response, ok := message.(Response); ok { + return s.Send(response) + } + return nil +} + +// Generated bidirectional handlers use typed Recv directly. +func (s *channelStreamServer[Request, Response]) RecvMsg(any) error { return nil } + +func (c *channelStreamClient[Request, Response]) AddCleanup(fn func()) { c.cleanup.add(fn) } +func (c *channelStreamClient[Request, Response]) Send(request Request) error { + c.sendMu.Lock() + defer c.sendMu.Unlock() + if c.closed { + return io.EOF + } + select { + case <-c.done: + return io.EOF + case <-c.ctx.Done(): + return c.ctx.Err() + default: + } + select { + case c.requests <- request: + return nil + case <-c.done: + return io.EOF + case <-c.ctx.Done(): + return c.ctx.Err() + } +} +func (c *channelStreamClient[Request, Response]) Recv() (Response, error) { + var zero Response + select { + case response, ok := <-c.responses: + if ok { + return response, nil + } + c.cleanup.finish() + if c.err != nil { + return zero, c.err + } + return zero, io.EOF + case <-c.ctx.Done(): + c.cleanup.finish() + return zero, c.ctx.Err() + } +} + +// Half-close only the request side; the server may still produce responses. +func (c *channelStreamClient[Request, Response]) CloseSend() error { + c.sendMu.Lock() + defer c.sendMu.Unlock() + if !c.closed { + c.closed = true + close(c.requests) + } + return nil +} diff --git a/pkg/grpc/channel_stream_test.go b/pkg/grpc/channel_stream_test.go new file mode 100644 index 000000000000..e5d8ca8ccc3e --- /dev/null +++ b/pkg/grpc/channel_stream_test.go @@ -0,0 +1,110 @@ +// SPDX-License-Identifier: MIT +package grpc + +import ( + "context" + "errors" + "io" + "sync" + "sync/atomic" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("In-process bidirectional channel streams", func() { + It("drains responses after half-close, then preserves the terminal error", func() { + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + terminal := errors.New("backend failed after tail") + serverDone := make(chan struct{}) + client := newChannelStream(ctx, 2, func(server *channelStreamServer[int, string]) error { + defer close(serverDone) + request, err := server.Recv() + if err != nil { + return err + } + if request != 7 { + return errors.New("wrong request") + } + _, err = server.Recv() + if err != io.EOF { + return errors.New("expected half-close") + } + if err = server.Send("tail"); err != nil { + return err + } + return terminal + }) + var cleaned atomic.Int32 + client.AddCleanup(func() { cleaned.Add(1) }) + Expect(client.Send(7)).To(Succeed()) + Expect(client.CloseSend()).To(Succeed()) + Expect(client.CloseSend()).To(Succeed()) + Eventually(serverDone).Should(BeClosed()) + Expect(cleaned.Load()).To(BeZero()) + tail, err := client.Recv() + Expect(err).NotTo(HaveOccurred()) + Expect(tail).To(Equal("tail")) + for range 2 { + _, err = client.Recv() + Expect(err).To(MatchError(terminal)) + } + Expect(cleaned.Load()).To(Equal(int32(1))) + Expect(client.Send(8)).To(MatchError(io.EOF)) + }) + It("unblocks pending sends when the server finishes without reading", func() { + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + stop := make(chan struct{}) + client := newChannelStream(ctx, 0, func(*channelStreamServer[int, int]) error { <-stop; return errors.New("stopped") }) + sent := make(chan error, 1) + go func() { sent <- client.Send(1) }() + close(stop) + Eventually(sent).Should(Receive(MatchError(io.EOF))) + _, err := client.Recv() + Expect(err).To(MatchError("stopped")) + }) + It("cancels blocked readers and writers and releases cleanup exactly once", func() { + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + serverDone := make(chan struct{}) + client := newChannelStream(ctx, 0, func(server *channelStreamServer[int, int]) error { + defer close(serverDone) + <-ctx.Done() + return ctx.Err() + }) + var cleaned atomic.Int32 + client.AddCleanup(func() { cleaned.Add(1) }) + sent := make(chan error, 1) + received := make(chan error, 1) + go func() { sent <- client.Send(1) }() + go func() { _, err := client.Recv(); received <- err }() + cancel() + Eventually(sent).Should(Receive(HaveOccurred())) + Eventually(received).Should(Receive(MatchError(context.Canceled))) + Eventually(serverDone).Should(BeClosed()) + Eventually(cleaned.Load).Should(Equal(int32(1))) + }) + It("allows concurrent half-closes without closing responses", func() { + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + client := newChannelStream(ctx, 1, func(server *channelStreamServer[int, int]) error { + _, err := server.Recv() + if err != io.EOF { + return err + } + return server.Send(42) + }) + var wg sync.WaitGroup + for range 8 { + wg.Go(func() { _ = client.CloseSend() }) + } + wg.Wait() + result, err := client.Recv() + Expect(err).NotTo(HaveOccurred()) + Expect(result).To(Equal(42)) + _, err = client.Recv() + Expect(err).To(MatchError(io.EOF)) + }) +}) diff --git a/pkg/grpc/embed.go b/pkg/grpc/embed.go index 58aab0861dc0..59b8b3f4a826 100644 --- a/pkg/grpc/embed.go +++ b/pkg/grpc/embed.go @@ -2,8 +2,6 @@ package grpc import ( "context" - "io" - "sync" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "google.golang.org/grpc" @@ -186,107 +184,27 @@ func (e *embedBackend) AudioTransform(ctx context.Context, in *pb.AudioTransform } func (e *embedBackend) AudioTransformStream(ctx context.Context, opts ...grpc.CallOption) (AudioTransformStreamClient, error) { - // In-process bidi stream is two channels paired with two facades: - // the server side reads requests / writes responses; the client side - // is its mirror. - reqs := make(chan *pb.AudioTransformFrameRequest, 4) - resps := make(chan *pb.AudioTransformFrameResponse, 4) - srvDone := make(chan error, 1) - - server := &embedBackendAudioTransformStream{ - ctx: ctx, - reqs: reqs, - resps: resps, - } - - go func() { - err := e.s.AudioTransformStream(server) - // Backend has finished — no more responses will arrive. - close(resps) - srvDone <- err - }() - - return &embedBackendAudioTransformStreamClient{ - ctx: ctx, - reqs: reqs, - resps: resps, - srvDone: srvDone, - cleanup: newStreamCleanup(ctx, nil), - }, nil + return newChannelStream(ctx, 4, func(stream *channelStreamServer[*pb.AudioTransformFrameRequest, *pb.AudioTransformFrameResponse]) error { + return e.s.AudioTransformStream(stream) + }), nil } func (e *embedBackend) AudioTranscriptionLive(ctx context.Context, opts ...grpc.CallOption) (AudioTranscriptionLiveClient, error) { - reqs := make(chan *pb.TranscriptLiveRequest, 4) - resps := make(chan *pb.TranscriptLiveResponse, 4) - srvDone := make(chan error, 1) - - server := &embedBackendAudioTranscriptionLiveStream{ - ctx: ctx, - reqs: reqs, - resps: resps, - } - - go func() { - err := e.s.AudioTranscriptionLive(server) - // Stash the terminal error BEFORE closing resps: a caller blocked in - // Recv wakes on the close and must find the error (the ready-ack - // contract surfaces Unimplemented through that first Recv). - srvDone <- err - close(resps) - }() - - return &embedBackendAudioTranscriptionLiveStreamClient{ - ctx: ctx, - reqs: reqs, - resps: resps, - srvDone: srvDone, - }, nil + return newChannelStream(ctx, 4, func(stream *channelStreamServer[*pb.TranscriptLiveRequest, *pb.TranscriptLiveResponse]) error { + return e.s.AudioTranscriptionLive(stream) + }), nil } func (e *embedBackend) Forward(ctx context.Context, opts ...grpc.CallOption) (ForwardClient, error) { - reqs := make(chan *pb.ForwardRequest, 8) - resps := make(chan *pb.ForwardReply, 8) - srvDone := make(chan error, 1) - - server := &embedBackendForwardStream{ctx: ctx, reqs: reqs, resps: resps} - - go func() { - err := e.s.Forward(server) - close(resps) - srvDone <- err - }() - - return &embedBackendForwardStreamClient{ - ctx: ctx, - reqs: reqs, - resps: resps, - srvDone: srvDone, - }, nil + return newChannelStream(ctx, 8, func(stream *channelStreamServer[*pb.ForwardRequest, *pb.ForwardReply]) error { + return e.s.Forward(stream) + }), nil } func (e *embedBackend) AudioToAudioStream(ctx context.Context, opts ...grpc.CallOption) (AudioToAudioStreamClient, error) { - reqs := make(chan *pb.AudioToAudioRequest, 8) - resps := make(chan *pb.AudioToAudioResponse, 8) - srvDone := make(chan error, 1) - - server := &embedBackendAudioToAudioStream{ - ctx: ctx, - reqs: reqs, - resps: resps, - } - - go func() { - err := e.s.AudioToAudioStream(server) - close(resps) - srvDone <- err - }() - - return &embedBackendAudioToAudioStreamClient{ - ctx: ctx, - reqs: reqs, - resps: resps, - srvDone: srvDone, - }, nil + return newChannelStream(ctx, 8, func(stream *channelStreamServer[*pb.AudioToAudioRequest, *pb.AudioToAudioResponse]) error { + return e.s.AudioToAudioStream(stream) + }), nil } func (e *embedBackend) ModelMetadata(ctx context.Context, in *pb.ModelOptions, opts ...grpc.CallOption) (*pb.ModelMetadataResponse, error) { @@ -342,302 +260,6 @@ func (e *embedBackend) Free(ctx context.Context) error { return err } -var _ pb.Backend_AudioTransformStreamServer = new(embedBackendAudioTransformStream) -var _ AudioTransformStreamClient = new(embedBackendAudioTransformStreamClient) -var _ pb.Backend_AudioToAudioStreamServer = new(embedBackendAudioToAudioStream) -var _ AudioToAudioStreamClient = new(embedBackendAudioToAudioStreamClient) -var _ pb.Backend_AudioTranscriptionLiveServer = new(embedBackendAudioTranscriptionLiveStream) -var _ AudioTranscriptionLiveClient = new(embedBackendAudioTranscriptionLiveStreamClient) - -// embedBackendAudioTransformStream is the server side of an in-process bidi -// stream. The hosted server reads requests from `reqs` (closed by client when -// done sending) and writes responses to `resps`. -type embedBackendAudioTransformStream struct { - ctx context.Context - reqs <-chan *pb.AudioTransformFrameRequest - resps chan<- *pb.AudioTransformFrameResponse -} - -func (e *embedBackendAudioTransformStream) Send(resp *pb.AudioTransformFrameResponse) error { - select { - case e.resps <- resp: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioTransformStream) Recv() (*pb.AudioTransformFrameRequest, error) { - select { - case req, ok := <-e.reqs: - if !ok { - return nil, io.EOF - } - return req, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioTransformStream) SetHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioTransformStream) SendHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioTransformStream) SetTrailer(md metadata.MD) {} -func (e *embedBackendAudioTransformStream) Context() context.Context { return e.ctx } -func (e *embedBackendAudioTransformStream) SendMsg(m any) error { - if x, ok := m.(*pb.AudioTransformFrameResponse); ok { - return e.Send(x) - } - return nil -} -func (e *embedBackendAudioTransformStream) RecvMsg(m any) error { - // gRPC bidi streaming uses Recv() directly; RecvMsg is unused on this path. - return nil -} - -// embedBackendAudioTransformStreamClient is the caller-facing side. It -// mirrors the server-side stream over the same channels. -type embedBackendAudioTransformStreamClient struct { - ctx context.Context - reqs chan<- *pb.AudioTransformFrameRequest - resps <-chan *pb.AudioTransformFrameResponse - srvDone <-chan error - closeOnce bool - cleanup *streamCleanup -} - -func (e *embedBackendAudioTransformStreamClient) AddCleanup(fn func()) { e.cleanup.add(fn) } - -func (e *embedBackendAudioTransformStreamClient) Send(req *pb.AudioTransformFrameRequest) error { - select { - case e.reqs <- req: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioTransformStreamClient) Recv() (*pb.AudioTransformFrameResponse, error) { - select { - case resp, ok := <-e.resps: - if !ok { - e.cleanup.finish() - // Server-side finished. Surface its terminal error if any. - select { - case err := <-e.srvDone: - if err != nil { - return nil, err - } - default: - } - return nil, io.EOF - } - return resp, nil - case <-e.ctx.Done(): - e.cleanup.finish() - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioTransformStreamClient) CloseSend() error { - if e.closeOnce { - return nil - } - e.closeOnce = true - close(e.reqs) - return nil -} - -func (e *embedBackendAudioTransformStreamClient) Context() context.Context { return e.ctx } - -// embedBackendAudioTranscriptionLiveStream is the in-process server-side -// handle for the bidirectional live ASR RPC. Mirrors -// embedBackendAudioTransformStream — the hosted server reads requests from -// `reqs` (closed by client when done sending) and writes responses to `resps`. -type embedBackendAudioTranscriptionLiveStream struct { - ctx context.Context - reqs <-chan *pb.TranscriptLiveRequest - resps chan<- *pb.TranscriptLiveResponse -} - -func (e *embedBackendAudioTranscriptionLiveStream) Send(resp *pb.TranscriptLiveResponse) error { - select { - case e.resps <- resp: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioTranscriptionLiveStream) Recv() (*pb.TranscriptLiveRequest, error) { - select { - case req, ok := <-e.reqs: - if !ok { - return nil, io.EOF - } - return req, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioTranscriptionLiveStream) SetHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioTranscriptionLiveStream) SendHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioTranscriptionLiveStream) SetTrailer(md metadata.MD) {} -func (e *embedBackendAudioTranscriptionLiveStream) Context() context.Context { return e.ctx } -func (e *embedBackendAudioTranscriptionLiveStream) SendMsg(m any) error { - if x, ok := m.(*pb.TranscriptLiveResponse); ok { - return e.Send(x) - } - return nil -} -func (e *embedBackendAudioTranscriptionLiveStream) RecvMsg(m any) error { - // gRPC bidi streaming uses Recv() directly; RecvMsg is unused on this path. - return nil -} - -// embedBackendAudioTranscriptionLiveStreamClient is the caller-facing side. -// It mirrors the server-side stream over the same channels. -type embedBackendAudioTranscriptionLiveStreamClient struct { - ctx context.Context - reqs chan<- *pb.TranscriptLiveRequest - resps <-chan *pb.TranscriptLiveResponse - srvDone <-chan error - closeOnce bool -} - -func (e *embedBackendAudioTranscriptionLiveStreamClient) Send(req *pb.TranscriptLiveRequest) error { - select { - case e.reqs <- req: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioTranscriptionLiveStreamClient) Recv() (*pb.TranscriptLiveResponse, error) { - select { - case resp, ok := <-e.resps: - if !ok { - // Server-side finished. Surface its terminal error if any. - select { - case err := <-e.srvDone: - if err != nil { - return nil, err - } - default: - } - return nil, io.EOF - } - return resp, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioTranscriptionLiveStreamClient) CloseSend() error { - if e.closeOnce { - return nil - } - e.closeOnce = true - close(e.reqs) - return nil -} - -func (e *embedBackendAudioTranscriptionLiveStreamClient) Context() context.Context { return e.ctx } - -// embedBackendAudioToAudioStream is the in-process server-side handle for -// the bidirectional any-to-any audio RPC. Mirrors embedBackendAudioTransform -// Stream — the hosted server reads requests from `reqs` (closed by client -// when done sending) and writes responses to `resps`. -type embedBackendAudioToAudioStream struct { - ctx context.Context - reqs <-chan *pb.AudioToAudioRequest - resps chan<- *pb.AudioToAudioResponse -} - -func (e *embedBackendAudioToAudioStream) Send(resp *pb.AudioToAudioResponse) error { - select { - case e.resps <- resp: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioToAudioStream) Recv() (*pb.AudioToAudioRequest, error) { - select { - case req, ok := <-e.reqs: - if !ok { - return nil, io.EOF - } - return req, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioToAudioStream) SetHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioToAudioStream) SendHeader(md metadata.MD) error { return nil } -func (e *embedBackendAudioToAudioStream) SetTrailer(md metadata.MD) {} -func (e *embedBackendAudioToAudioStream) Context() context.Context { return e.ctx } -func (e *embedBackendAudioToAudioStream) SendMsg(m any) error { - if x, ok := m.(*pb.AudioToAudioResponse); ok { - return e.Send(x) - } - return nil -} -func (e *embedBackendAudioToAudioStream) RecvMsg(m any) error { return nil } - -type embedBackendAudioToAudioStreamClient struct { - ctx context.Context - reqs chan<- *pb.AudioToAudioRequest - resps <-chan *pb.AudioToAudioResponse - srvDone <-chan error - closeOnce bool -} - -func (e *embedBackendAudioToAudioStreamClient) Send(req *pb.AudioToAudioRequest) error { - select { - case e.reqs <- req: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendAudioToAudioStreamClient) Recv() (*pb.AudioToAudioResponse, error) { - select { - case resp, ok := <-e.resps: - if !ok { - // Server goroutine writes to srvDone immediately after closing - // resps; block (cap with ctx) so we don't race past a real error. - select { - case err := <-e.srvDone: - if err != nil { - return nil, err - } - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } - return nil, io.EOF - } - return resp, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendAudioToAudioStreamClient) CloseSend() error { - if e.closeOnce { - return nil - } - e.closeOnce = true - close(e.reqs) - return nil -} - -func (e *embedBackendAudioToAudioStreamClient) Context() context.Context { return e.ctx } - var _ pb.Backend_AudioTranscriptionStreamServer = new(embedBackendAudioTranscriptionStream) type embedBackendAudioTranscriptionStream struct { @@ -787,94 +409,3 @@ func (e *embedBackendServerStream) SendMsg(m any) error { func (e *embedBackendServerStream) RecvMsg(m any) error { return nil } - -var _ pb.Backend_ForwardServer = new(embedBackendForwardStream) -var _ ForwardClient = new(embedBackendForwardStreamClient) - -// embedBackendForwardStream is the server-side handle for an in-process -// Forward bidi stream. The hosted backend reads requests from `reqs` -// (closed by the client when done sending) and writes replies to -// `resps`. -type embedBackendForwardStream struct { - ctx context.Context - reqs <-chan *pb.ForwardRequest - resps chan<- *pb.ForwardReply -} - -func (e *embedBackendForwardStream) Send(resp *pb.ForwardReply) error { - select { - case e.resps <- resp: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendForwardStream) Recv() (*pb.ForwardRequest, error) { - select { - case req, ok := <-e.reqs: - if !ok { - return nil, io.EOF - } - return req, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendForwardStream) SetHeader(md metadata.MD) error { return nil } -func (e *embedBackendForwardStream) SendHeader(md metadata.MD) error { return nil } -func (e *embedBackendForwardStream) SetTrailer(md metadata.MD) {} -func (e *embedBackendForwardStream) Context() context.Context { return e.ctx } -func (e *embedBackendForwardStream) SendMsg(m any) error { - if x, ok := m.(*pb.ForwardReply); ok { - return e.Send(x) - } - return nil -} -func (e *embedBackendForwardStream) RecvMsg(m any) error { return nil } - -// embedBackendForwardStreamClient is the caller-facing side. Mirrors -// the server-side stream over the same channels. -type embedBackendForwardStreamClient struct { - ctx context.Context - reqs chan<- *pb.ForwardRequest - resps <-chan *pb.ForwardReply - srvDone <-chan error - once sync.Once -} - -func (e *embedBackendForwardStreamClient) Send(req *pb.ForwardRequest) error { - select { - case e.reqs <- req: - return nil - case <-e.ctx.Done(): - return e.ctx.Err() - } -} - -func (e *embedBackendForwardStreamClient) Recv() (*pb.ForwardReply, error) { - select { - case resp, ok := <-e.resps: - if !ok { - select { - case err := <-e.srvDone: - if err != nil { - return nil, err - } - default: - } - return nil, io.EOF - } - return resp, nil - case <-e.ctx.Done(): - return nil, e.ctx.Err() - } -} - -func (e *embedBackendForwardStreamClient) CloseSend() error { - e.once.Do(func() { close(e.reqs) }) - return nil -} - -func (e *embedBackendForwardStreamClient) Context() context.Context { return e.ctx } diff --git a/pkg/grpc/interface.go b/pkg/grpc/interface.go index c4d6d4967950..0e15734b42a9 100644 --- a/pkg/grpc/interface.go +++ b/pkg/grpc/interface.go @@ -13,6 +13,7 @@ type AnimationMetadataModel interface { } type AIModel interface { + MotionStream(context.Context, func() (*pb.MotionRequest, error), func(*pb.MotionResponse) error) error Busy() bool Lock() Unlock() diff --git a/pkg/grpc/motion.go b/pkg/grpc/motion.go new file mode 100644 index 000000000000..c154add591fa --- /dev/null +++ b/pkg/grpc/motion.go @@ -0,0 +1,101 @@ +// SPDX-License-Identifier: MIT +package grpc + +import ( + "context" + "sync" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +type MotionStreamClient interface { + Send(*pb.MotionRequest) error + Recv() (*pb.MotionResponse, error) + CloseSend() error + Context() context.Context +} + +type motionStreamClient struct { + pb.Backend_MotionStreamClient + closeOnce sync.Once + closer func() +} + +// Keep the connection alive after CloseSend until final output is drained. +func (s *motionStreamClient) Recv() (*pb.MotionResponse, error) { + resp, err := s.Backend_MotionStreamClient.Recv() + if err != nil { + s.release() + } + return resp, err +} + +func (s *motionStreamClient) release() { + s.closeOnce.Do(func() { + if s.closer != nil { + s.closer() + } + }) +} + +// MotionStream holds the watchdog busy marker for the session lifetime. +func (c *Client) MotionStream(ctx context.Context, opts ...grpc.CallOption) (MotionStreamClient, error) { + if !c.parallel { + if !c.opMutex.TryLock() { + return nil, status.Error(codes.ResourceExhausted, "backend is busy") + } + } + c.setBusy(true) + completeRequest := c.wdMark() + + cleanup := func() { + completeRequest() + c.setBusy(false) + if !c.parallel { + c.opMutex.Unlock() + } + } + + conn, err := c.dial() + if err != nil { + cleanup() + return nil, err + } + client := pb.NewBackendClient(conn) + stream, err := client.MotionStream(ctx, opts...) + if err != nil { + _ = conn.Close() + cleanup() + return nil, err + } + result := &motionStreamClient{ + Backend_MotionStreamClient: stream, + closer: func() { + _ = conn.Close() + cleanup() + }, + } + go func() { <-stream.Context().Done(); result.release() }() + return result, nil +} + +func (s *server) MotionStream(stream pb.Backend_MotionStreamServer) error { + first, err := stream.Recv() + if err != nil { + return err + } + if err := s.checkModelIdentity(first); err != nil { + return err + } + pending := true + return s.llm.MotionStream(stream.Context(), func() (*pb.MotionRequest, error) { + if pending { + pending = false + return first, nil + } + return stream.Recv() + }, stream.Send) +} diff --git a/pkg/grpc/motion_embed.go b/pkg/grpc/motion_embed.go new file mode 100644 index 000000000000..99768b55e820 --- /dev/null +++ b/pkg/grpc/motion_embed.go @@ -0,0 +1,14 @@ +// SPDX-License-Identifier: MIT +package grpc + +import ( + "context" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "google.golang.org/grpc" +) + +func (e *embedBackend) MotionStream(ctx context.Context, opts ...grpc.CallOption) (MotionStreamClient, error) { + return newChannelStream(ctx, 4, func(stream *channelStreamServer[*pb.MotionRequest, *pb.MotionResponse]) error { + return e.s.MotionStream(stream) + }), nil +} diff --git a/pkg/grpc/motion_test.go b/pkg/grpc/motion_test.go new file mode 100644 index 000000000000..5bc874b0edb8 --- /dev/null +++ b/pkg/grpc/motion_test.go @@ -0,0 +1,66 @@ +// SPDX-License-Identifier: MIT +package grpc + +import ( + "context" + "github.com/mudler/LocalAI/pkg/grpc/base" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "io" + "time" +) + +type motionEcho struct{ base.Base } + +func (*motionEcho) Load(*pb.ModelOptions) error { return nil } +func (*motionEcho) MotionStream(_ context.Context, recv func() (*pb.MotionRequest, error), send func(*pb.MotionResponse) error) error { + for { + r, err := recv() + if err == io.EOF { + return nil + } + if err != nil { + return err + } + if err := send(&pb.MotionResponse{Output: r.Input}); err != nil { + return err + } + } +} + +var _ = Describe("MotionStream RPC", func() { + It("round-trips binary payloads, drains half-close and cancels idle receivers", func() { + Provide("test://motion-echo", &motionEcho{}) + c := NewClient("test://motion-echo", true, nil, false) + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + stream, err := c.MotionStream(ctx) + Expect(err).NotTo(HaveOccurred()) + Expect(stream.Send(&pb.MotionRequest{Input: []byte{8, 128, 1}})).To(Succeed()) + reply, err := stream.Recv() + Expect(err).NotTo(HaveOccurred()) + Expect(reply.Output).To(Equal([]byte{8, 128, 1})) + Expect(stream.CloseSend()).To(Succeed()) + _, err = stream.Recv() + Expect(err).To(Equal(io.EOF)) + other, err := c.MotionStream(ctx) + Expect(err).NotTo(HaveOccurred()) + cancel() + done := make(chan error, 1) + go func() { _, err := other.Recv(); done <- err }() + Eventually(done, time.Second).Should(Receive(HaveOccurred())) + }) + It("rejects stale model identity before invoking native code", func() { + s := &server{llm: &motionEcho{}, loadedIdentity: "actual"} + e := &embedBackend{s: s} + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + stream, err := e.MotionStream(ctx) + Expect(err).NotTo(HaveOccurred()) + Expect(stream.Send(&pb.MotionRequest{ModelIdentity: "stale"})).To(Succeed()) + _, err = stream.Recv() + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("actual")) + }) +}) diff --git a/pkg/model/remote_shutdown_test.go b/pkg/model/remote_shutdown_test.go index cc010c45cc18..05a5a0a819d4 100644 --- a/pkg/model/remote_shutdown_test.go +++ b/pkg/model/remote_shutdown_test.go @@ -3,6 +3,7 @@ package model_test import ( "context" "errors" + "os/exec" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" @@ -83,9 +84,11 @@ var _ = Describe("ShutdownModel in distributed mode", func() { unloader.unloadErr = remoteErr modelLoader.SetRemoteUnloader(unloader) + sleepPath, err := exec.LookPath("sleep") + Expect(err).NotTo(HaveOccurred()) localProcess := process.New( process.WithTemporaryStateDir(), - process.WithName("/bin/sleep"), + process.WithName(sleepPath), process.WithArgs("300"), ) Expect(localProcess.Run()).To(Succeed()) @@ -95,7 +98,7 @@ var _ = Describe("ShutdownModel in distributed mode", func() { } }) - _, err := modelLoader.LoadModel("mixed", "mixed", func(_, _, _ string) (*model.Model, error) { + _, err = modelLoader.LoadModel("mixed", "mixed", func(_, _, _ string) (*model.Model, error) { return model.NewModel("mixed", "local", localProcess), nil }) Expect(err).NotTo(HaveOccurred()) diff --git a/pkg/motion/frame_queue.go b/pkg/motion/frame_queue.go new file mode 100644 index 000000000000..cf260555e92c --- /dev/null +++ b/pkg/motion/frame_queue.go @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: MIT +package motion + +import ( + "math" + "time" +) + +const ( + UploadWindow = 8 + UploadMaxBytes = 8 << 20 + UploadFrameMaxBytes = 2 << 20 + UploadMaxQueueAge = 250 * time.Millisecond +) + +// QueuedFrame owns its encoded Input until inference starts or it is discarded. +type QueuedFrame struct { + Sequence uint64 + SourceTime int64 + Data []byte + Arrived time.Time +} + +// FrameQueue has one external lock/owner. It never owns the active inference. +type FrameQueue struct { + frames []QueuedFrame + bytes int + Dropped uint64 +} + +func (q *FrameQueue) OldestAge(now time.Time) time.Duration { + if len(q.frames) == 0 { + return 0 + } + return max(0, now.Sub(q.frames[0].Arrived)) +} + +func (q *FrameQueue) Len() int { return len(q.frames) } +func (q *FrameQueue) Bytes() int { return q.bytes } + +// Thin uniformly samples the entire pending span, including both endpoints. +// It releases removed buffers immediately rather than retaining slice capacity. +func (q *FrameQueue) thin(keep int) []uint64 { + survivors := make([]QueuedFrame, 0, keep) + released := make([]uint64, 0, len(q.frames)-keep) + selected := make(map[int]bool, keep) + if keep == 1 { + selected[len(q.frames)-1] = true + } else if keep > 1 { + selected[0], selected[len(q.frames)-1] = true, true + first, last := q.frames[0].SourceTime, q.frames[len(q.frames)-1].SourceTime + previous := 0 + for sample := 1; sample < keep-1; sample++ { + target := float64(first) + float64(last-first)*float64(sample)/float64(keep-1) + nearest := previous + 1 + for candidate := nearest + 1; candidate <= len(q.frames)-(keep-sample); candidate++ { + if math.Abs(float64(q.frames[candidate].SourceTime)-target) < math.Abs(float64(q.frames[nearest].SourceTime)-target) { + nearest = candidate + } + } + selected[nearest] = true + previous = nearest + } + } + bytes := 0 + for i, frame := range q.frames { + if selected[i] { + survivors = append(survivors, frame) + bytes += len(frame.Data) + } else { + released = append(released, frame.Sequence) + } + } + q.frames, q.bytes = survivors, bytes + q.Dropped += uint64(len(released)) + return released +} + +func (q *FrameQueue) Expire(now time.Time) []uint64 { + var released []uint64 + for len(q.frames) > 0 && now.Sub(q.frames[0].Arrived) >= UploadMaxQueueAge { + f := q.frames[0] + released = append(released, f.Sequence) + q.bytes -= len(f.Data) + q.frames[0] = QueuedFrame{} + q.frames = q.frames[1:] + } + q.Dropped += uint64(len(released)) + return released +} + +// Offer accepts one bounded frame, then thins before another read can grow the queue. +func (q *FrameQueue) Offer(frame QueuedFrame) (released []uint64, pressure bool) { + released = q.Expire(frame.Arrived) + q.frames = append(q.frames, frame) + q.bytes += len(frame.Data) + pressure = len(q.frames) >= 6 || q.bytes >= 6<<20 + if pressure { + released = append(released, q.thin(max(1, len(q.frames)/2))...) + } + return released, pressure +} + +func (q *FrameQueue) Take() (QueuedFrame, bool) { + if len(q.frames) == 0 { + return QueuedFrame{}, false + } + f := q.frames[0] + q.frames[0] = QueuedFrame{} + q.frames = q.frames[1:] + q.bytes -= len(f.Data) + return f, true +} diff --git a/pkg/motion/frame_queue_test.go b/pkg/motion/frame_queue_test.go new file mode 100644 index 000000000000..d7372b0001e7 --- /dev/null +++ b/pkg/motion/frame_queue_test.go @@ -0,0 +1,73 @@ +// SPDX-License-Identifier: MIT +package motion + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "testing" + "time" +) + +func TestFrameQueue(t *testing.T) { RegisterFailHandler(Fail); RunSpecs(t, "Motion frame queue") } + +var _ = Describe("Bounded temporal frame queue", func() { + It("thins across the pending span, preserving endpoints and earlier ties", func() { + q := &FrameQueue{} + now := time.Now() + var released []uint64 + + for i := 1; i <= 6; i++ { + var pressure bool + released, pressure = q.Offer(QueuedFrame{Sequence: uint64(i), SourceTime: int64(i), Data: []byte{1}, Arrived: now}) + Expect(pressure).To(Equal(i == 6)) + } + + Expect(released).To(Equal([]uint64{2, 4, 5})) + var kept []uint64 + + for { + f, ok := q.Take() + if !ok { + break + } + kept = append(kept, f.Sequence) + } + + Expect(kept).To(Equal([]uint64{1, 3, 6})) + }) + It("uses timestamp spacing instead of index spacing for irregular captures", func() { + q := &FrameQueue{} + + for i, stamp := range []int64{0, 1, 2, 50, 99, 100} { + q.Offer(QueuedFrame{Sequence: uint64(i + 1), SourceTime: stamp, Data: []byte{1}, Arrived: time.Now()}) + } + var kept []uint64 + + for { + f, ok := q.Take() + if !ok { + break + } + kept = append(kept, f.Sequence) + } + + Expect(kept).To(Equal([]uint64{1, 4, 6})) + }) + It("bounds bytes under repeated large uploads and expires by server arrival time", func() { + q := &FrameQueue{} + now := time.Now() + data := make([]byte, UploadFrameMaxBytes) + + for i := 1; i <= 100; i++ { + q.Offer(QueuedFrame{Sequence: uint64(i), Data: data, Arrived: now}) + Expect(q.Len()).To(BeNumerically("<=", UploadWindow)) + Expect(q.Bytes()).To(BeNumerically("<", 6<<20)) + } + + Expect(q.Expire(now.Add(UploadMaxQueueAge))).NotTo(BeEmpty()) + + Expect(q.Len()).To(BeZero()) + Expect(q.Bytes()).To(BeZero()) + Expect(q.Dropped).To(Equal(uint64(100))) + }) +}) diff --git a/pkg/motion/proto/motion.proto b/pkg/motion/proto/motion.proto new file mode 100644 index 000000000000..bbc5529072b7 --- /dev/null +++ b/pkg/motion/proto/motion.proto @@ -0,0 +1,75 @@ +// SPDX-License-Identifier: MIT +syntax = "proto3"; +package localai.motion.v1; +option go_package = "github.com/mudler/LocalAI/pkg/motion/proto;motionpb"; + +// Public protocol, independent of LocalAI's internal backend service. +// All float arrays are joint-major, component-minor packed float32. +message Input { + oneof payload { + Frame frame = 1; + bool reset_state = 2; + } +} +message Frame { + uint64 sequence = 1; + int64 source_time_us = 2; // Nonnegative, relative to caller-declared origin. + bytes rgb = 3; // Tightly packed RGB8, width*height*3 bytes. + uint32 width = 4; + uint32 height = 5; + repeated float subject_box = 6; // Optional source-pixel XYXY, exactly 4 values. + uint64 subject_id = 7; +} +message Definition { + string schema = 1; + repeated string joint_names = 2; + repeated sint32 parents = 3; + uint32 root = 4; + repeated float rest_local_translations = 5; + repeated float rest_local_rotations = 6; + // Coordinate/shape/timing conventions only; never model paths. + map conventions = 7; + string profile = 8; + repeated string channels = 9; +} +message Pose { + uint64 sequence = 1; + int64 source_time_us = 2; + uint64 epoch = 3; + uint64 track_epoch = 4; + uint32 flags = 5; // 1 reset, 2 reused crop, 4 full image, 8 detector, 16 caller box. + repeated float positions = 6; + repeated float local_rotations = 7; // SOMA XYZW; absent for SMPL. + repeated float local_translations = 8; + repeated float anchor = 9; // SMPL WXYZ; absent for SOMA. + repeated float root_axis_angle = 10; // SOMA radians; absent for SMPL. + repeated float root_translation = 11; // Zero; no world trajectory recovery. + repeated float box = 12; + // Optional source-image pixel XY pairs in Definition.joint_names order. + // Upper-left origin, +X right, +Y down; no mirroring or crop transform. + // Same submitted Frame sequence/time and dimensions. Empty if not projectable. + repeated float image_positions = 13; + // Optional SMPL-only closed-interval displacement, metres in the END pose's + // anchor-local basis. Rotate by anchor for gravity-aligned Z-up displacement. + // Not velocity. Interval is [displacement_start_time_us, source_time_us]. + // Empty when unsupported. Consumers must fence epochs and missing intervals; + // no absolute world position or contact-refined trajectory is implied. + repeated float root_displacement = 14; + int64 displacement_start_time_us = 15; +} +message Event { + string type = 1; // warmup, lost, ambiguous, reset, error, closed + string message = 2; + uint64 sequence = 3; + int64 source_time_us = 4; + uint64 epoch = 5; + uint64 track_epoch = 6; + uint32 flags = 7; +} +message Output { + oneof payload { + Definition definition = 1; + Pose pose = 2; + Event event = 3; + } +} diff --git a/pkg/motion/validation.go b/pkg/motion/validation.go new file mode 100644 index 000000000000..15f394fe8578 --- /dev/null +++ b/pkg/motion/validation.go @@ -0,0 +1,36 @@ +// SPDX-License-Identifier: MIT +package motion + +import ( + "fmt" + motionpb "github.com/mudler/LocalAI/pkg/motion/proto" + "math" +) + +// ValidateFrame bounds allocation and rejects invalid subject boxes before a +// host forwards data to a native inference process. +func ValidateFrame(f *motionpb.Frame) error { + if f == nil { + return fmt.Errorf("frame or reset required") + } + n := uint64(f.Width) * uint64(f.Height) + if f.Width < 8 || f.Height < 8 || f.Width > 32766 || f.Height > 32766 || n > 16000000 || uint64(len(f.Rgb)) != n*3 || f.SourceTimeUs < 0 { + return fmt.Errorf("invalid RGB dimensions, byte length or timestamp") + } + b := f.SubjectBox + if len(b) == 0 { + return nil + } + if len(b) != 4 { + return fmt.Errorf("subject_box requires four coordinates") + } + for _, v := range b { + if math.IsNaN(float64(v)) || math.IsInf(float64(v), 0) { + return fmt.Errorf("subject_box must be finite") + } + } + if b[0] < 0 || b[1] < 0 || b[2] > float32(f.Width-1) || b[3] > float32(f.Height-1) || b[2] <= b[0] || b[3] <= b[1] { + return fmt.Errorf("subject_box must be ordered XYXY within image pixel coordinates") + } + return nil +} diff --git a/scripts/lib/backend-filter.mjs b/scripts/lib/backend-filter.mjs index 9312cbf6d0c0..069433d3087a 100644 --- a/scripts/lib/backend-filter.mjs +++ b/scripts/lib/backend-filter.mjs @@ -345,6 +345,12 @@ export function protoChangeIsAdditive(previousText, currentText) { // first-match-wins per changed file, so specific entries must precede the // scripts/build/ catch-all at the bottom. export const SHARED_BUILD_INPUTS = [ + { + // GEM-X links the public motion protocol and shared frame validation. + matches: file => file.startsWith("pkg/motion/") && !file.endsWith("_test.go"), + linux: item => item.backend === "gemxcpp", + darwin: item => item.backend === "gemxcpp", + }, { // Every language consumes backend.proto: Dockerfile.python COPYs it, Go // backends regenerate their stubs from it via `make protogen-go`, the C++ diff --git a/scripts/lib/backend-filter_test.mjs b/scripts/lib/backend-filter_test.mjs index 6999c921a127..25d942ae642d 100644 --- a/scripts/lib/backend-filter_test.mjs +++ b/scripts/lib/backend-filter_test.mjs @@ -10,6 +10,7 @@ import assert from "node:assert/strict"; import { filterMatrix, + SHARED_BUILD_INPUTS, inferBackendPath, inferBackendPathDarwin, } from "./backend-filter.mjs"; @@ -19,6 +20,11 @@ test("kimodocpp maps to its native Go wrapper on Linux and Darwin", () => { assert.equal(inferBackendPathDarwin({ backend: "kimodocpp", lang: "go" }), "backend/go/kimodocpp/"); }); +test("gemxcpp maps to its native Go wrapper on Linux and Darwin", () => { + assert.equal(inferBackendPath({ backend: "gemxcpp", dockerfile: "./backend/Dockerfile.golang" }), "backend/go/gemxcpp/"); + assert.equal(inferBackendPathDarwin({ backend: "gemxcpp", lang: "go" }), "backend/go/gemxcpp/"); +}); + test("trellis2cpp maps to its Go backend source directory", () => { assert.equal( inferBackendPath({ @@ -596,3 +602,13 @@ test("unresolvable proto revisions conservatively rebuild everything", () => { assert.equal(filtered.length, includes.length); assert.equal(filteredDarwin.length, includesDarwin.length); }); + + +test("public motion schema and validation changes rebuild GEM-X on both OSes", () => { + for (const file of ["pkg/motion/proto/motion.proto", "pkg/motion/validation.go"]) { + const rules = SHARED_BUILD_INPUTS.filter(rule => rule.matches(file)); + assert.ok(rules.some(rule => rule.linux({ backend: "gemxcpp", dockerfile: "./backend/Dockerfile.golang" }))); + assert.ok(rules.some(rule => rule.darwin({ backend: "gemxcpp", lang: "go" }))); + assert.ok(!rules.some(rule => rule.linux({ backend: "kimodocpp", dockerfile: "./backend/Dockerfile.golang" }))); + } +}); diff --git a/swagger/docs.go b/swagger/docs.go index 063447074be4..86b328506052 100644 --- a/swagger/docs.go +++ b/swagger/docs.go @@ -1467,6 +1467,146 @@ const docTemplate = `{ } } }, + "/api/motion/sessions": { + "post": { + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "motion" + ], + "summary": "Create a motion capture session", + "parameters": [ + { + "description": "Model, output profile and source clock origin", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/localai.MotionSessionRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "$ref": "#/definitions/localai.MotionSessionInfo" + } + } + } + } + }, + "/api/motion/sessions/{id}": { + "get": { + "tags": [ + "motion" + ], + "summary": "Get motion session status", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/localai.MotionSessionInfo" + } + } + } + }, + "delete": { + "tags": [ + "motion" + ], + "summary": "Close a motion session", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "No Content" + } + } + } + }, + "/api/motion/sessions/{id}/poses": { + "get": { + "tags": [ + "motion" + ], + "summary": "Stream motion frames and poses over a duplex WebSocket", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "101": { + "description": "Switching Protocols" + } + } + } + }, + "/api/motion/sessions/{id}/tickets": { + "post": { + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "motion" + ], + "summary": "Issue a motion WebSocket ticket (30 seconds, single use)", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + }, + { + "description": "Browser origin", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/localai.MotionTicketRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "$ref": "#/definitions/auth.WebSocketTicketResponse" + } + } + } + } + }, "/api/nodes/models": { "get": { "tags": [ @@ -4446,6 +4586,17 @@ const docTemplate = `{ } }, "definitions": { + "auth.WebSocketTicketResponse": { + "type": "object", + "properties": { + "expires_at": { + "type": "string" + }, + "ticket": { + "type": "string" + } + } + }, "config.Gallery": { "type": "object", "properties": { @@ -5233,6 +5384,57 @@ const docTemplate = `{ } } }, + "localai.MotionSessionInfo": { + "type": "object", + "properties": { + "frames": { + "type": "integer" + }, + "id": { + "type": "string" + }, + "model": { + "type": "string" + }, + "poses": { + "type": "integer" + }, + "profile": { + "type": "string" + }, + "time_origin": { + "type": "string" + }, + "upload_accepted": { + "type": "integer" + }, + "upload_dropped": { + "type": "integer" + } + } + }, + "localai.MotionSessionRequest": { + "type": "object", + "properties": { + "model": { + "type": "string" + }, + "profile": { + "type": "string" + }, + "time_origin": { + "type": "string" + } + } + }, + "localai.MotionTicketRequest": { + "type": "object", + "properties": { + "origin": { + "type": "string" + } + } + }, "localai.TTSModelVoices": { "type": "object", "properties": { diff --git a/swagger/swagger.json b/swagger/swagger.json index 7ccfd6984462..f0a0ab8b9cf7 100644 --- a/swagger/swagger.json +++ b/swagger/swagger.json @@ -1464,6 +1464,146 @@ } } }, + "/api/motion/sessions": { + "post": { + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "motion" + ], + "summary": "Create a motion capture session", + "parameters": [ + { + "description": "Model, output profile and source clock origin", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/localai.MotionSessionRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "$ref": "#/definitions/localai.MotionSessionInfo" + } + } + } + } + }, + "/api/motion/sessions/{id}": { + "get": { + "tags": [ + "motion" + ], + "summary": "Get motion session status", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/localai.MotionSessionInfo" + } + } + } + }, + "delete": { + "tags": [ + "motion" + ], + "summary": "Close a motion session", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "No Content" + } + } + } + }, + "/api/motion/sessions/{id}/poses": { + "get": { + "tags": [ + "motion" + ], + "summary": "Stream motion frames and poses over a duplex WebSocket", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "101": { + "description": "Switching Protocols" + } + } + } + }, + "/api/motion/sessions/{id}/tickets": { + "post": { + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "motion" + ], + "summary": "Issue a motion WebSocket ticket (30 seconds, single use)", + "parameters": [ + { + "type": "string", + "description": "Session ID", + "name": "id", + "in": "path", + "required": true + }, + { + "description": "Browser origin", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/localai.MotionTicketRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "$ref": "#/definitions/auth.WebSocketTicketResponse" + } + } + } + } + }, "/api/nodes/models": { "get": { "tags": [ @@ -4443,6 +4583,17 @@ } }, "definitions": { + "auth.WebSocketTicketResponse": { + "type": "object", + "properties": { + "expires_at": { + "type": "string" + }, + "ticket": { + "type": "string" + } + } + }, "config.Gallery": { "type": "object", "properties": { @@ -5230,6 +5381,57 @@ } } }, + "localai.MotionSessionInfo": { + "type": "object", + "properties": { + "frames": { + "type": "integer" + }, + "id": { + "type": "string" + }, + "model": { + "type": "string" + }, + "poses": { + "type": "integer" + }, + "profile": { + "type": "string" + }, + "time_origin": { + "type": "string" + }, + "upload_accepted": { + "type": "integer" + }, + "upload_dropped": { + "type": "integer" + } + } + }, + "localai.MotionSessionRequest": { + "type": "object", + "properties": { + "model": { + "type": "string" + }, + "profile": { + "type": "string" + }, + "time_origin": { + "type": "string" + } + } + }, + "localai.MotionTicketRequest": { + "type": "object", + "properties": { + "origin": { + "type": "string" + } + } + }, "localai.TTSModelVoices": { "type": "object", "properties": { diff --git a/swagger/swagger.yaml b/swagger/swagger.yaml index 14224a6b499d..c675018da86b 100644 --- a/swagger/swagger.yaml +++ b/swagger/swagger.yaml @@ -1,5 +1,12 @@ basePath: / definitions: + auth.WebSocketTicketResponse: + properties: + expires_at: + type: string + ticket: + type: string + type: object config.Gallery: properties: artifact_verification: @@ -579,6 +586,39 @@ definitions: success: type: boolean type: object + localai.MotionSessionInfo: + properties: + frames: + type: integer + id: + type: string + model: + type: string + poses: + type: integer + profile: + type: string + time_origin: + type: string + upload_accepted: + type: integer + upload_dropped: + type: integer + type: object + localai.MotionSessionRequest: + properties: + model: + type: string + profile: + type: string + time_origin: + type: string + type: object + localai.MotionTicketRequest: + properties: + origin: + type: string + type: object localai.TTSModelVoices: properties: model: @@ -4551,6 +4591,96 @@ paths: summary: Estimate VRAM usage for a model tags: - config + /api/motion/sessions: + post: + consumes: + - application/json + parameters: + - description: Model, output profile and source clock origin + in: body + name: request + required: true + schema: + $ref: '#/definitions/localai.MotionSessionRequest' + produces: + - application/json + responses: + "201": + description: Created + schema: + $ref: '#/definitions/localai.MotionSessionInfo' + summary: Create a motion capture session + tags: + - motion + /api/motion/sessions/{id}: + delete: + parameters: + - description: Session ID + in: path + name: id + required: true + type: string + responses: + "204": + description: No Content + summary: Close a motion session + tags: + - motion + get: + parameters: + - description: Session ID + in: path + name: id + required: true + type: string + responses: + "200": + description: OK + schema: + $ref: '#/definitions/localai.MotionSessionInfo' + summary: Get motion session status + tags: + - motion + /api/motion/sessions/{id}/poses: + get: + parameters: + - description: Session ID + in: path + name: id + required: true + type: string + responses: + "101": + description: Switching Protocols + summary: Stream motion frames and poses over a duplex WebSocket + tags: + - motion + /api/motion/sessions/{id}/tickets: + post: + consumes: + - application/json + parameters: + - description: Session ID + in: path + name: id + required: true + type: string + - description: Browser origin + in: body + name: request + required: true + schema: + $ref: '#/definitions/localai.MotionTicketRequest' + produces: + - application/json + responses: + "201": + description: Created + schema: + $ref: '#/definitions/auth.WebSocketTicketResponse' + summary: Issue a motion WebSocket ticket (30 seconds, single use) + tags: + - motion /api/nodes/{id}/max-replicas-per-model: delete: parameters: