ONNX
security
malware-detection
Vigil / source /pkg /compactonnx /loader.go
turentomer's picture
Make published source self-contained
a2a3348 verified
Raw
History Blame Contribute Delete
11.2 kB
//go:build !noonnx
package compactonnx
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"os"
"path/filepath"
"runtime"
"strconv"
"sync"
"huggingface.co/turenlabs/Vigil/source/pkg/compactmodel"
ort "github.com/yalue/onnxruntime_go"
)
var compactRuntime struct {
sync.Mutex
initialized bool
libraryPath string
version string
}
// Model owns one immutable in-memory ONNX session and its bound tensors.
type Model struct {
metadata Metadata
markers Markers
session *ort.AdvancedSession
input *ort.Tensor[float32]
output *ort.Tensor[float32]
mu sync.Mutex
closed bool
}
// Load validates the adjacent pair, initializes ONNX Runtime, validates the
// graph interface, and creates a session directly from the verified bytes.
func Load(config Config) (*Model, error) {
pair, err := readArtifactPair(config.ModelPath, config.MetadataPath)
if err != nil {
return nil, err
}
metadata, err := decodeMetadata(pair.metadata)
if err != nil {
return nil, err
}
if err := validateMetadata(metadata, pair); err != nil {
return nil, err
}
runtimeVersion, err := initializeRuntime(config.RuntimeLibrary)
if err != nil {
return nil, err
}
if err := validateGraph(pair.model, metadata); err != nil {
return nil, err
}
input, err := ort.NewEmptyTensor[float32](ort.NewShape(1, compactmodel.TotalFeatures))
if err != nil {
return nil, fmt.Errorf("create compact ONNX input: %w", err)
}
output, err := ort.NewEmptyTensor[float32](ort.NewShape(1, 1))
if err != nil {
input.Destroy()
return nil, fmt.Errorf("create compact ONNX output: %w", err)
}
session, err := ort.NewAdvancedSessionWithONNXData(
pair.model,
[]string{InputName},
[]string{OutputName},
[]ort.Value{input},
[]ort.Value{output},
nil,
)
if err != nil {
output.Destroy()
input.Destroy()
return nil, fmt.Errorf("create compact ONNX session: %w", err)
}
return &Model{
metadata: metadata,
markers: Markers{
ModelLoaded: true,
ModelType: "compact_hashed_linear_onnx",
ModelSHA256: pair.modelSHA256,
ModelSizeBytes: int64(len(pair.model)),
MetadataSHA256: pair.metadataSHA256,
PreprocessingContractSHA256: compactmodel.ContractSHA256(),
Runtime: "onnxruntime",
RuntimeVersion: runtimeVersion,
FallbackUsed: false,
},
session: session,
input: input,
output: output,
}, nil
}
func decodeMetadata(data []byte) (Metadata, error) {
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
var metadata Metadata
if err := decoder.Decode(&metadata); err != nil {
return Metadata{}, fmt.Errorf("decode compact metadata: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
if err == nil {
return Metadata{}, errors.New("decode compact metadata: trailing JSON value")
}
return Metadata{}, fmt.Errorf("decode compact metadata trailer: %w", err)
}
return metadata, nil
}
func validateMetadata(metadata Metadata, pair artifactPair) error {
if metadata.SchemaVersion != MetadataSchema {
return fmt.Errorf("compact metadata schema mismatch: got %q", metadata.SchemaVersion)
}
if metadata.Model.SHA256 != pair.modelSHA256 {
return errors.New("compact metadata model SHA-256 mismatch")
}
if metadata.Model.SizeBytes != int64(len(pair.model)) {
return errors.New("compact metadata model size mismatch")
}
if metadata.Model.SizeBytes <= 0 || metadata.Model.SizeBytes > MaxModelBytes {
return errors.New("compact metadata model size is outside the release limit")
}
if metadata.Model.Family != "hashed-word-char-linear" {
return errors.New("compact metadata model family mismatch")
}
if metadata.Model.InputName != InputName || metadata.Model.InputFeatures != compactmodel.TotalFeatures {
return errors.New("compact metadata input contract mismatch")
}
if metadata.Model.OutputName != OutputName || metadata.Model.OutputElements != 1 {
return errors.New("compact metadata output contract mismatch")
}
if err := compactmodel.ValidateMetadata(metadata.preprocessingMetadata()); err != nil {
return err
}
if math.IsNaN(metadata.Threshold) || math.IsInf(metadata.Threshold, 0) || metadata.Threshold <= 0 || metadata.Threshold >= 1 {
return errors.New("compact metadata threshold must be finite and between zero and one")
}
return nil
}
func validateGraph(data []byte, metadata Metadata) error {
inputs, outputs, err := ort.GetInputOutputInfoWithONNXData(data)
if err != nil {
return fmt.Errorf("inspect compact ONNX graph: %w", err)
}
if len(inputs) != 1 || !validTensor(inputs[0], InputName, compactmodel.TotalFeatures) {
return fmt.Errorf("compact ONNX input mismatch: want one float tensor %s [batch,%d]", InputName, compactmodel.TotalFeatures)
}
if len(outputs) != 1 || !validTensor(outputs[0], OutputName, 1) {
return fmt.Errorf("compact ONNX output mismatch: want one float tensor %s [batch,1]", OutputName)
}
modelMetadata, err := ort.GetModelMetadataWithONNXData(data)
if err != nil {
return fmt.Errorf("inspect compact ONNX metadata: %w", err)
}
defer modelMetadata.Destroy()
required := map[string]string{
"model_family": "hashed-word-char-linear",
"word_features": strconv.Itoa(compactmodel.WordFeatures),
"char_features": strconv.Itoa(compactmodel.CharFeatures),
"structured_features": strconv.Itoa(compactmodel.StructuredFeatures),
"preprocessing": "sklearn-murmurhash3 word(1,2)+char(4)+package-structure-v1",
}
for key, want := range required {
got, found, err := modelMetadata.LookupCustomMetadataMap(key)
if err != nil || !found || got != want {
return fmt.Errorf("compact ONNX internal metadata mismatch for %s", key)
}
}
thresholdText, found, err := modelMetadata.LookupCustomMetadataMap("threshold")
if err != nil || !found {
return errors.New("compact ONNX internal threshold metadata is missing")
}
threshold, err := strconv.ParseFloat(thresholdText, 64)
if err != nil || threshold != metadata.Threshold {
return errors.New("compact ONNX internal threshold metadata mismatch")
}
return nil
}
func validTensor(info ort.InputOutputInfo, name string, width int) bool {
if info.Name != name || info.OrtValueType != ort.ONNXTypeTensor || info.DataType != ort.TensorElementDataTypeFloat {
return false
}
return len(info.Dimensions) == 2 &&
(info.Dimensions[0] == -1 || info.Dimensions[0] == 1) &&
info.Dimensions[1] == int64(width)
}
func initializeRuntime(explicit string) (string, error) {
compactRuntime.Lock()
defer compactRuntime.Unlock()
if compactRuntime.initialized {
if explicit != "" {
resolved, err := canonicalRuntimeLibrary(explicit)
if err != nil {
return "", err
}
if resolved != compactRuntime.libraryPath {
return "", errors.New("ONNX Runtime is already initialized from a different library")
}
}
return compactRuntime.version, nil
}
if ort.IsInitialized() {
return "", errors.New("ONNX Runtime was initialized outside the compact loader; exact runtime provenance is unavailable")
}
library := explicit
if library == "" {
library = os.Getenv("ONNXRUNTIME_LIB")
}
if library == "" {
library = discoverRuntimeLibrary()
}
if library == "" {
return "", errors.New("ONNX Runtime library is required; compact scoring has no fallback")
}
resolved, err := canonicalRuntimeLibrary(library)
if err != nil {
return "", err
}
ort.SetSharedLibraryPath(resolved)
if err := ort.InitializeEnvironment(); err != nil {
return "", fmt.Errorf("initialize ONNX Runtime: %w", err)
}
compactRuntime.initialized = true
compactRuntime.libraryPath = resolved
compactRuntime.version = ort.GetVersion()
return compactRuntime.version, nil
}
func canonicalRuntimeLibrary(path string) (string, error) {
if hasTraversal(path) {
return "", errors.New("ONNX Runtime path contains traversal")
}
abs, err := filepath.Abs(path)
if err != nil {
return "", fmt.Errorf("resolve ONNX Runtime path: %w", err)
}
resolved, err := filepath.EvalSymlinks(abs)
if err != nil {
return "", fmt.Errorf("resolve ONNX Runtime library: %w", err)
}
info, err := os.Lstat(resolved)
if err != nil {
return "", fmt.Errorf("inspect ONNX Runtime library: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return "", errors.New("ONNX Runtime library is not a regular file")
}
return resolved, nil
}
func discoverRuntimeLibrary() string {
name := "libonnxruntime.so"
switch runtime.GOOS {
case "darwin":
name = "libonnxruntime.dylib"
case "windows":
name = "onnxruntime.dll"
}
var candidates []string
if executable, err := os.Executable(); err == nil {
directory := filepath.Dir(executable)
candidates = append(candidates, filepath.Join(directory, name), filepath.Join(directory, "lib", name))
}
if working, err := os.Getwd(); err == nil {
candidates = append(candidates, filepath.Join(working, name), filepath.Join(working, "lib", name))
}
switch runtime.GOOS {
case "darwin":
candidates = append(candidates, filepath.Join("/opt/homebrew/lib", name), filepath.Join("/usr/local/lib", name))
case "linux":
candidates = append(candidates, filepath.Join("/usr/lib", name), filepath.Join("/usr/local/lib", name), filepath.Join("/usr/lib/x86_64-linux-gnu", name), filepath.Join("/usr/lib/aarch64-linux-gnu", name))
}
for _, candidate := range candidates {
if _, err := os.Stat(candidate); err == nil {
return candidate
}
}
return ""
}
// Score vectorizes and scores one whole package. There is no fallback branch.
func (model *Model) Score(pkg compactmodel.Package) (float64, error) {
vector, err := compactmodel.Vectorize(pkg)
if err != nil {
return 0, err
}
model.mu.Lock()
defer model.mu.Unlock()
if model.closed || model.session == nil {
return 0, errors.New("compact ONNX model is closed")
}
copy(model.input.GetData(), vector)
if err := model.session.Run(); err != nil {
return 0, fmt.Errorf("run compact ONNX: %w", err)
}
values := model.output.GetData()
if len(values) != 1 {
return 0, fmt.Errorf("compact ONNX output length mismatch: got %d", len(values))
}
probability := float64(values[0])
if math.IsNaN(probability) || math.IsInf(probability, 0) || probability < 0 || probability > 1 {
return 0, errors.New("compact ONNX output is not a finite probability")
}
return probability, nil
}
func (model *Model) Threshold() float64 { return model.metadata.Threshold }
func (model *Model) Markers() Markers { return model.markers }
// Close releases the session and bound tensors. The process-global runtime is
// intentionally retained because other loaded sessions may still use it.
func (model *Model) Close() error {
model.mu.Lock()
defer model.mu.Unlock()
if model.closed {
return nil
}
model.closed = true
var errs []error
if model.session != nil {
errs = append(errs, model.session.Destroy())
}
if model.output != nil {
errs = append(errs, model.output.Destroy())
}
if model.input != nil {
errs = append(errs, model.input.Destroy())
}
return errors.Join(errs...)
}