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: