Compare commits

...

23 Commits

Author SHA1 Message Date
dorcaslitunya 3447495fcd Uncomment out kernel version 2025-02-10 14:23:48 +00:00
dorcaslitunya b3257dc9a6 Modify go-sev-guest version 2025-02-10 11:45:48 +00:00
dorcaslitunya d408c7bef2 Formatting changes 2025-02-10 11:39:18 +00:00
dorcaslitunya 417f0c7291 Add kernel changes 2025-02-10 11:32:02 +00:00
dorcaslitunya f1f7a89a6c Modify buildroot config to enable vTPM attestations 2025-02-10 11:20:02 +00:00
dependabot[bot] 132bfdf76a NOISSUE - Bump the go-dependency group across 1 directory with 10 updates (#366)
CI / ci (push) Has been cancelled
Bumps the go-dependency group with 5 updates in the / directory:

| Package | From | To |
| --- | --- | --- |
| [github.com/caarlos0/env/v11](https://github.com/caarlos0/env) | `11.2.2` | `11.3.1` |
| [github.com/google/go-sev-guest](https://github.com/google/go-sev-guest) | `0.11.1` | `0.12.1` |
| [github.com/spf13/pflag](https://github.com/spf13/pflag) | `1.0.5` | `1.0.6` |
| [go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc](https://github.com/open-telemetry/opentelemetry-go-contrib) | `0.57.0` | `0.59.0` |
| [github.com/docker/docker](https://github.com/docker/docker) | `27.4.0+incompatible` | `27.5.1+incompatible` |



Updates `github.com/caarlos0/env/v11` from 11.2.2 to 11.3.1
- [Release notes](https://github.com/caarlos0/env/releases)
- [Changelog](https://github.com/caarlos0/env/blob/main/.goreleaser.yml)
- [Commits](https://github.com/caarlos0/env/compare/v11.2.2...v11.3.1)

Updates `github.com/google/go-sev-guest` from 0.11.1 to 0.12.1
- [Release notes](https://github.com/google/go-sev-guest/releases)
- [Changelog](https://github.com/google/go-sev-guest/blob/main/.goreleaser.yaml)
- [Commits](https://github.com/google/go-sev-guest/compare/v0.11.1...v0.12.1)

Updates `github.com/spf13/pflag` from 1.0.5 to 1.0.6
- [Release notes](https://github.com/spf13/pflag/releases)
- [Commits](https://github.com/spf13/pflag/compare/v1.0.5...v1.0.6)

Updates `go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc` from 0.57.0 to 0.59.0
- [Release notes](https://github.com/open-telemetry/opentelemetry-go-contrib/releases)
- [Changelog](https://github.com/open-telemetry/opentelemetry-go-contrib/blob/main/CHANGELOG.md)
- [Commits](https://github.com/open-telemetry/opentelemetry-go-contrib/compare/zpages/v0.57.0...zpages/v0.59.0)

Updates `go.opentelemetry.io/otel/trace` from 1.32.0 to 1.34.0
- [Release notes](https://github.com/open-telemetry/opentelemetry-go/releases)
- [Changelog](https://github.com/open-telemetry/opentelemetry-go/blob/main/CHANGELOG.md)
- [Commits](https://github.com/open-telemetry/opentelemetry-go/compare/v1.32.0...v1.34.0)

Updates `golang.org/x/crypto` from 0.30.0 to 0.32.0
- [Commits](https://github.com/golang/crypto/compare/v0.30.0...v0.32.0)

Updates `google.golang.org/grpc` from 1.68.1 to 1.69.4
- [Release notes](https://github.com/grpc/grpc-go/releases)
- [Commits](https://github.com/grpc/grpc-go/compare/v1.68.1...v1.69.4)

Updates `google.golang.org/protobuf` from 1.35.2 to 1.36.3

Updates `github.com/docker/docker` from 27.4.0+incompatible to 27.5.1+incompatible
- [Release notes](https://github.com/docker/docker/releases)
- [Commits](https://github.com/docker/docker/compare/v27.4.0...v27.5.1)

Updates `golang.org/x/term` from 0.27.0 to 0.28.0
- [Commits](https://github.com/golang/term/compare/v0.27.0...v0.28.0)

---
updated-dependencies:
- dependency-name: github.com/caarlos0/env/v11
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: github.com/google/go-sev-guest
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: github.com/spf13/pflag
  dependency-type: direct:production
  update-type: version-update:semver-patch
  dependency-group: go-dependency
- dependency-name: go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: go.opentelemetry.io/otel/trace
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: golang.org/x/crypto
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: google.golang.org/grpc
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: google.golang.org/protobuf
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: github.com/docker/docker
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: golang.org/x/term
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2025-02-04 15:06:23 +01:00
Washington Kigani Kamadi 51f2a02e4a NOISSUE - Update env for new manager deployment (#367)
CI / ci (push) Waiting to run
* fix env

Signed-off-by: WashingtonKK <washingtonkigan@gmail.com>

* update kernel and rootfs location

Signed-off-by: WashingtonKK <washingtonkigan@gmail.com>

* update manager host

Signed-off-by: WashingtonKK <washingtonkigan@gmail.com>

---------

Signed-off-by: WashingtonKK <washingtonkigan@gmail.com>
2025-02-04 14:25:02 +01:00
Smith Jilks da88fe1e45 COCOS-346 - Explore cloud init for Cloud setup (#357)
CI / ci (push) Has been cancelled
Rust CI Pipeline / rust-check (push) Has been cancelled
* Add qemu cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Update qemu cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Add qemu cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Update qemu cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Update qemu cloud config

* Update cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Update cloud init

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Add cloud init README.md

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

* Add cocos release workflow

Signed-off-by: Jilks Smith <smithjilks@gmail.com>

---------

Signed-off-by: Jilks Smith <smithjilks@gmail.com>
2025-01-31 15:48:26 +01:00
dependabot[bot] 5969ae3bcb NOISSUE - Update SEV requirement (#330)
Updates the requirements on [sev](https://github.com/virtee/sev) to permit the latest version.

Updates `sev` to 5.0.0
- [Commits](https://github.com/virtee/sev/compare/v4.0.0...v5.0.0)

---
updated-dependencies:
- dependency-name: sev
  dependency-type: direct:production
  dependency-group: rs-dependencies
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2025-01-31 15:47:12 +01:00
Sammy Kerata Oina b5c65f6c3f Update agent CVM gRPC certificate keys for consistency (#361)
CI / ci (push) Has been cancelled
Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2025-01-29 12:25:21 +01:00
Washington Kigani Kamadi 5bc7eb2c8a Add manager service client mocks (#359)
CI / ci (push) Has been cancelled
Signed-off-by: WashingtonKK <washingtonkigan@gmail.com>
2025-01-27 09:49:25 +01:00
Sammy Kerata Oina 58b401e0de Update dependency for sev-snp-measure-go to latest version (#358)
CI / checkproto (push) Has been cancelled
CI / ci (push) Has been cancelled
Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2025-01-21 15:19:41 +01:00
Sammy Kerata Oina 881aaaab0f NOISSUE - Set env automatically (#355)
* new agent structure

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix lint

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* cvm tests fix

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* add cli and test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* restore result cli

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* pass certs and env

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* update go

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* downgrade

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* downgrade again

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* simplify

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* simplify

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* configure cvms

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* remove unused gRPC API files and server implementation

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* refactor: use constants for CLI command flags and environment variables

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

---------

Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2025-01-20 13:46:18 +01:00
Sammy Kerata Oina 1f32f516b0 NOISSUE - Simplify manager to vm provision only (#353)
* new agent structure

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix lint

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* cvm tests fix

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* manager server, for vm provisioning

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix lint

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* add cli and test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* restore result cli

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix failing tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix failing test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* refactor: remove context from docker struct and use local context in Run method

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* delete: remove unused gRPC API and related server implementation

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

---------

Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2025-01-20 11:56:18 +01:00
Sammy Kerata Oina ecad6514f3 COCOS-344 - New agent structure (#350)
CI / checkproto (push) Has been cancelled
CI / ci (push) Has been cancelled
* new agent structure

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* minor fixes and testing

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix lint

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* cvm tests fix

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix cli test

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* rename

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* rename cvm to cvms plural

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* rename service

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* remove context

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* refactor: reorder parameters in NewAlgorithm functions and update CVMClient to CVMSClient

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix(tests): update SendEvent mock to include an additional parameter

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* move expectations

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix(tests): move event initialization to the correct scope in service tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

* fix(tests): update SendEvent mock to use EXPECT instead of On in service tests

Signed-off-by: Sammy Oina <sammyoina@gmail.com>

---------

Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2025-01-17 12:50:53 +01:00
Drasko DRASKOVIC 59b8057e5c Update README.md (#348) 2024-12-31 01:12:01 +01:00
dorcaslitunya 961f8025ca Update README.md (#341) 2024-12-16 15:54:45 +01:00
Sammy Kerata Oina 35c09be0d9 fix test (#335)
Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2024-12-16 15:52:59 +01:00
dependabot[bot] 4c49be5684 Bump the go-dependency group across 1 directory with 7 updates (#331)
Bumps the go-dependency group with 4 updates in the / directory: [github.com/absmach/magistrala](https://github.com/absmach/magistrala), [golang.org/x/crypto](https://github.com/golang/crypto), [google.golang.org/grpc](https://github.com/grpc/grpc-go) and [github.com/docker/docker](https://github.com/docker/docker).


Updates `github.com/absmach/magistrala` from 0.14.1-0.20240709113739-04c359462746 to 0.15.1
- [Release notes](https://github.com/absmach/magistrala/releases)
- [Commits](https://github.com/absmach/magistrala/commits/v0.15.1)

Updates `github.com/stretchr/testify` from 1.9.0 to 1.10.0
- [Release notes](https://github.com/stretchr/testify/releases)
- [Commits](https://github.com/stretchr/testify/compare/v1.9.0...v1.10.0)

Updates `golang.org/x/crypto` from 0.29.0 to 0.30.0
- [Commits](https://github.com/golang/crypto/compare/v0.29.0...v0.30.0)

Updates `golang.org/x/sync` from 0.9.0 to 0.10.0
- [Commits](https://github.com/golang/sync/compare/v0.9.0...v0.10.0)

Updates `google.golang.org/grpc` from 1.68.0 to 1.68.1
- [Release notes](https://github.com/grpc/grpc-go/releases)
- [Commits](https://github.com/grpc/grpc-go/compare/v1.68.0...v1.68.1)

Updates `github.com/docker/docker` from 27.3.1+incompatible to 27.4.0+incompatible
- [Release notes](https://github.com/docker/docker/releases)
- [Commits](https://github.com/docker/docker/compare/v27.3.1...v27.4.0)

Updates `golang.org/x/term` from 0.26.0 to 0.27.0
- [Commits](https://github.com/golang/term/compare/v0.26.0...v0.27.0)

---
updated-dependencies:
- dependency-name: github.com/absmach/magistrala
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: github.com/stretchr/testify
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: golang.org/x/crypto
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: golang.org/x/sync
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: google.golang.org/grpc
  dependency-type: direct:production
  update-type: version-update:semver-patch
  dependency-group: go-dependency
- dependency-name: github.com/docker/docker
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
- dependency-name: golang.org/x/term
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: go-dependency
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2024-12-16 12:37:30 +01:00
Sammy Kerata Oina 3cd64546f3 return vm config (#334)
Signed-off-by: Sammy Oina <sammyoina@gmail.com>
2024-12-11 16:39:01 +01:00
Danko Miladinovic e48f184075 NOISSUE - Add launch TCB info to VM info (#333)
* add launch TCB to VM info

* add mutex for AP

* add policy info to run test

* fix manager Run test

* add SEV-SNP check
2024-12-11 15:53:42 +01:00
Dušan Borovčanin 0315e7ddfa Merge pull request #332 from danko-miladinovic/atls 2024-12-11 12:20:06 +01:00
danko-miladinovic 394a73cef3 fix close notify messages 2024-12-10 15:56:55 +00:00
103 changed files with 5019 additions and 5659 deletions
+2 -2
View File
@@ -33,8 +33,8 @@ jobs:
- name: Set up protoc
run: |
PROTOC_VERSION=28.1
PROTOC_GEN_VERSION=v1.34.2
PROTOC_VERSION=29.0
PROTOC_GEN_VERSION=v1.36.0
PROTOC_GRPC_VERSION=v1.5.1
# Download and install protoc
+15 -8
View File
@@ -1,9 +1,9 @@
name: Build and Release
name: Build and Release Hal
on:
push:
tags:
- '*'
- "*"
jobs:
build:
@@ -32,8 +32,8 @@ jobs:
with:
root-reserve-mb: 35000
swap-size-mb: 1024
remove-dotnet: 'true'
remove-android: 'true'
remove-dotnet: "true"
remove-android: "true"
- name: Check free space
run: |
echo "Free space:"
@@ -48,26 +48,33 @@ jobs:
- name: Checkout cocos
uses: actions/checkout@v4
with:
repository: 'ultravioletrs/cocos'
repository: "ultravioletrs/cocos"
path: cocos
- name: Checkout buildroot
uses: actions/checkout@v4
with:
repository: 'buildroot/buildroot'
repository: "buildroot/buildroot"
path: buildroot
ref: 2024.11-rc2
- name: Build
- name: Build hal
run: |
cd buildroot
make BR2_EXTERNAL=../cocos/hal/linux cocos_defconfig
make
- name: Build cocos
run: |
cd cocos
make
- name: Release
uses: softprops/action-gh-release@v2
with:
files: |
buildroot/output/images/bzImage
buildroot/output/images/rootfs.cpio.gz
cocos/build/cocos-agent
cocos/build/cocos-cli
cocos/build/cocos-manager
+1
View File
@@ -37,6 +37,7 @@ protoc:
protoc -I. --go_out=. --go_opt=paths=source_relative --go-grpc_out=. --go-grpc_opt=paths=source_relative agent/agent.proto
protoc -I. --go_out=. --go_opt=paths=source_relative --go-grpc_out=. --go-grpc_opt=paths=source_relative manager/manager.proto
protoc -I. --go_out=. --go_opt=paths=source_relative --go-grpc_out=. --go-grpc_opt=paths=source_relative agent/events/events.proto
protoc -I. --go_out=. --go_opt=paths=source_relative --go-grpc_out=. --go-grpc_opt=paths=source_relative agent/cvms/cvms.proto
mocks:
mockery --config ./mockery.yml
+52 -39
View File
@@ -1,65 +1,78 @@
# Cocos AI
<div align="center">
# Cocos AI 🥥
**Confidential Computing System for AI**
**Made with ❤️ by [Ultraviolet](https://ultraviolet.rs/)**
[![codecov](https://codecov.io/gh/ultravioletrs/cocos/graph/badge.svg?token=HX01LR01K9)](https://codecov.io/gh/ultravioletrs/cocos)
![Go report card](https://goreportcard.com/badge/github.com/ultravioletrs/cocos)
[![Go report card](https://goreportcard.com/badge/github.com/ultravioletrs/cocos)](https://goreportcard.com/report/github.com/ultravioletrs/cocos)
[![License](https://img.shields.io/badge/license-Apache--2.0-blue)](LICENSE)
[Cocos AI (Confdential Computing System for AI/ML)][cocos] is a platform for secure multiparty computation (SMPC)
based on the [Confidential Computing][cc] and [Trusted Execution Environments (TEEs)][tee].
### [Guide](https://docs.cocos.ultraviolet.rs) | [Contributing](CONTRIBUTING.md) | [Website](https://cocos.ai/)
</div>
## Introduction 🚀
Cocos AI is a **cutting-edge platform** designed to enable secure multiparty computation (SMPC) using **Confidential Computing** and **Trusted Execution Environments (TEEs)**.
It empowers organizations to collaboratively process sensitive data for AI/ML workloads while ensuring:
- 🔒 **Data Privacy**: Your data stays encrypted and secure throughout the computation.
- 🛡️ **Trust and Integrity**: Protected by hardware enclaves with robust remote attestation protocols.
- 🤝 **Seamless Collaboration**: Multiple organizations can work together without exposing sensitive information.
<p align="center">
<img src="https://cocos.ai/images/Collaborative%20AI.drawio.svg" width="500" height="500">
<img src="https://cocos.ai/images/Collaborative%20AI.drawio.svg" alt="Cocos AI Illustration" width="400" height="400">
</p>
With Cocos AI it becomes possible to run AI/ML workloads on combined datasets from multiple organizations
while guaranteeing the privacy and security of the data and the algorithm.
Data is always encrypted, protected by hardware secure enclaves (Trusted Execution Environments),
attested via secure remote attestation protocols, and invisible to cloud processors or any other
3rd party to which computation is offloaded.
## Features 🛠️
## Features
Cocos AI provides essential features for secure and efficient collaborative AI/ML:
Cocos AI is implementing the following features:
- 🖥️ **TEE Enablement and Monitoring**: Secure VM management for deploying and monitoring workloads.
- 🛡️ **Hardware Abstraction Layer (HAL)**: Built on a hardened Linux kernel, secure bootloader, and minimal root filesystem (minimal TCB).
- 🕵️ **In-Enclave Agent and Networking Controller**: Essential system software for managing secure workloads.
- 🔒 **Encrypted Data Transfer**: Asynchronous data transfer and secure result delivery.
- 🛠️ **API for Platform Manipulation**: Programmatic control for managing workloads.
-**Attestation and Verification Tools**: Hardware- and software-supported attestation for integrity assurance.
- 🖱️ **Command-Line Interface (CLI)**: A user-friendly CLI for system interaction.
- TEE enablement, deployment and monitoring (secure VM manager)
- HAL for TEEs based on hardened Linux kernel, secure bootloader and custom-tailored embedded rootfs for minimal TCB
- In-enclave agent, netowrking controller and other system software
- Encrypted asynchronous data transfer and result delivery
- API for programmable platform manipulation
- HW and SW supported attestation with verification tools
- CLI for system interaction
## Usage
Clone the repo and create binaries:
## 🚀 Quick Start
### Clone the Repository and Build Binaries
```bash
git clone git@github.com:ultravioletrs/cocos.git
make
```
This will create 3 binaries:
This will generate three binaries:
```bash
ls build/
# cocos-agent cocos-cli cocos-manager
```
- Manager can be deployed on the AMD SEV-SNP host
- Agent can be built into [EOS][eos]-based HAL
- CLI can be used to communicate to remote Agent.
### Deployment Overview:
- **Manager**: Deploy on the AMD SEV-SNP host to orchestrate workloads.
- **Agent**: Build into the [EOS](https://github.com/ultravioletrs/eos)-based HAL for secure enclave management.
- **CLI**: Interact with remote agents to control operations.
## Documentation
## 📚 Documentation
Project documentation is hosted at [Cocos AI official docs page][docs].
Comprehensive documentation is available at the [official documentation page](https://docs.cocos.ultraviolet.rs).
For CLI usage details, visit the [CLI Documentation](https://docs.cocos.ultraviolet.rs/cli).
Documentation is generated from the [docs repository](https://github.com/ultravioletrs/docs).
Documentation is automatically generated from the [docs repository](https://github.com/ultravioletrs/docs). Contributions to documentation are welcome!
## License
Cocos AI is published under permissive open-source [Apache-2.0](LICENSE) license.
## 🛡️ License
[cc]: https://confidentialcomputing.io/white-papers-reports/
[cocos]: https://cocos.ai/
[rel]: https://github.com/ultravioletrs/cocos/releases
[tee]: https://en.wikipedia.org/wiki/Trusted_execution_environment
[docs]: https://docs.cocos.ultraviolet.rs
[cli]: https://docs.cocos.ultraviolet.rs/cli
[eos]: https://github.com/ultravioletrs/eos
Cocos AI is published under the permissive open-source [Apache-2.0](LICENSE) license. Contributions are encouraged and appreciated!
## 🌐 Links and Resources
- [Cocos AI Website](https://cocos.ai/)
- [Official Releases](https://github.com/ultravioletrs/cocos/releases)
- [Confidential Computing Overview](https://confidentialcomputing.io/white-papers-reports/)
- [Trusted Execution Environments (TEEs)](https://en.wikipedia.org/wiki/Trusted_execution_environment)
+57 -176
View File
@@ -3,8 +3,8 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.34.2
// protoc v5.28.1
// protoc-gen-go v1.36.0
// protoc v5.29.0
// source: agent/agent.proto
package agent
@@ -24,21 +24,18 @@ const (
)
type AlgoRequest struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
Algorithm []byte `protobuf:"bytes,1,opt,name=algorithm,proto3" json:"algorithm,omitempty"`
Requirements []byte `protobuf:"bytes,2,opt,name=requirements,proto3" json:"requirements,omitempty"`
unknownFields protoimpl.UnknownFields
Algorithm []byte `protobuf:"bytes,1,opt,name=algorithm,proto3" json:"algorithm,omitempty"`
Requirements []byte `protobuf:"bytes,2,opt,name=requirements,proto3" json:"requirements,omitempty"`
sizeCache protoimpl.SizeCache
}
func (x *AlgoRequest) Reset() {
*x = AlgoRequest{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AlgoRequest) String() string {
@@ -49,7 +46,7 @@ func (*AlgoRequest) ProtoMessage() {}
func (x *AlgoRequest) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[0]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -79,18 +76,16 @@ func (x *AlgoRequest) GetRequirements() []byte {
}
type AlgoResponse struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *AlgoResponse) Reset() {
*x = AlgoResponse{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AlgoResponse) String() string {
@@ -101,7 +96,7 @@ func (*AlgoResponse) ProtoMessage() {}
func (x *AlgoResponse) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[1]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -117,21 +112,18 @@ func (*AlgoResponse) Descriptor() ([]byte, []int) {
}
type DataRequest struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
Dataset []byte `protobuf:"bytes,1,opt,name=dataset,proto3" json:"dataset,omitempty"`
Filename string `protobuf:"bytes,2,opt,name=filename,proto3" json:"filename,omitempty"`
unknownFields protoimpl.UnknownFields
Dataset []byte `protobuf:"bytes,1,opt,name=dataset,proto3" json:"dataset,omitempty"`
Filename string `protobuf:"bytes,2,opt,name=filename,proto3" json:"filename,omitempty"`
sizeCache protoimpl.SizeCache
}
func (x *DataRequest) Reset() {
*x = DataRequest{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *DataRequest) String() string {
@@ -142,7 +134,7 @@ func (*DataRequest) ProtoMessage() {}
func (x *DataRequest) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[2]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -172,18 +164,16 @@ func (x *DataRequest) GetFilename() string {
}
type DataResponse struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *DataResponse) Reset() {
*x = DataResponse{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[3]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[3]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *DataResponse) String() string {
@@ -194,7 +184,7 @@ func (*DataResponse) ProtoMessage() {}
func (x *DataResponse) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[3]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -210,18 +200,16 @@ func (*DataResponse) Descriptor() ([]byte, []int) {
}
type ResultRequest struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ResultRequest) Reset() {
*x = ResultRequest{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[4]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[4]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ResultRequest) String() string {
@@ -232,7 +220,7 @@ func (*ResultRequest) ProtoMessage() {}
func (x *ResultRequest) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[4]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -248,20 +236,17 @@ func (*ResultRequest) Descriptor() ([]byte, []int) {
}
type ResultResponse struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
File []byte `protobuf:"bytes,1,opt,name=file,proto3" json:"file,omitempty"`
unknownFields protoimpl.UnknownFields
File []byte `protobuf:"bytes,1,opt,name=file,proto3" json:"file,omitempty"`
sizeCache protoimpl.SizeCache
}
func (x *ResultResponse) Reset() {
*x = ResultResponse{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[5]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[5]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ResultResponse) String() string {
@@ -272,7 +257,7 @@ func (*ResultResponse) ProtoMessage() {}
func (x *ResultResponse) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[5]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -295,20 +280,17 @@ func (x *ResultResponse) GetFile() []byte {
}
type AttestationRequest struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
ReportData []byte `protobuf:"bytes,1,opt,name=report_data,json=reportData,proto3" json:"report_data,omitempty"` // Should be of length 64.
unknownFields protoimpl.UnknownFields
ReportData []byte `protobuf:"bytes,1,opt,name=report_data,json=reportData,proto3" json:"report_data,omitempty"` // Should be of length 64.
sizeCache protoimpl.SizeCache
}
func (x *AttestationRequest) Reset() {
*x = AttestationRequest{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[6]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[6]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AttestationRequest) String() string {
@@ -319,7 +301,7 @@ func (*AttestationRequest) ProtoMessage() {}
func (x *AttestationRequest) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[6]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -342,20 +324,17 @@ func (x *AttestationRequest) GetReportData() []byte {
}
type AttestationResponse struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
File []byte `protobuf:"bytes,1,opt,name=file,proto3" json:"file,omitempty"`
unknownFields protoimpl.UnknownFields
File []byte `protobuf:"bytes,1,opt,name=file,proto3" json:"file,omitempty"`
sizeCache protoimpl.SizeCache
}
func (x *AttestationResponse) Reset() {
*x = AttestationResponse{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_agent_proto_msgTypes[7]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_agent_proto_msgTypes[7]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AttestationResponse) String() string {
@@ -366,7 +345,7 @@ func (*AttestationResponse) ProtoMessage() {}
func (x *AttestationResponse) ProtoReflect() protoreflect.Message {
mi := &file_agent_agent_proto_msgTypes[7]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -477,104 +456,6 @@ func file_agent_agent_proto_init() {
if File_agent_agent_proto != nil {
return
}
if !protoimpl.UnsafeEnabled {
file_agent_agent_proto_msgTypes[0].Exporter = func(v any, i int) any {
switch v := v.(*AlgoRequest); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[1].Exporter = func(v any, i int) any {
switch v := v.(*AlgoResponse); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[2].Exporter = func(v any, i int) any {
switch v := v.(*DataRequest); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[3].Exporter = func(v any, i int) any {
switch v := v.(*DataResponse); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[4].Exporter = func(v any, i int) any {
switch v := v.(*ResultRequest); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[5].Exporter = func(v any, i int) any {
switch v := v.(*ResultResponse); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[6].Exporter = func(v any, i int) any {
switch v := v.(*AttestationRequest); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_agent_proto_msgTypes[7].Exporter = func(v any, i int) any {
switch v := v.(*AttestationResponse); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
+1 -1
View File
@@ -4,7 +4,7 @@
// Code generated by protoc-gen-go-grpc. DO NOT EDIT.
// versions:
// - protoc-gen-go-grpc v1.5.1
// - protoc v5.28.1
// - protoc v5.29.0
// source: agent/agent.proto
package agent
+3
View File
@@ -46,4 +46,7 @@ func AlgorithmArgsFromContext(ctx context.Context) []string {
type Algorithm interface {
// Run executes the algorithm and returns the result.
Run() error
// Stop stops the algorithm.
Stop() error
}
+24 -7
View File
@@ -20,29 +20,46 @@ type binary struct {
stderr io.Writer
stdout io.Writer
args []string
cmd *exec.Cmd
}
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string, args []string) algorithm.Algorithm {
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string, args []string, cmpID string) algorithm.Algorithm {
return &binary{
algoFile: algoFile,
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc},
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc, CmpID: cmpID},
stdout: &logging.Stdout{Logger: logger},
args: args,
}
}
func (b *binary) Run() error {
cmd := exec.Command(b.algoFile, b.args...)
cmd.Stderr = b.stderr
cmd.Stdout = b.stdout
b.cmd = exec.Command(b.algoFile, b.args...)
b.cmd.Stderr = b.stderr
b.cmd.Stdout = b.stdout
if err := cmd.Start(); err != nil {
if err := b.cmd.Start(); err != nil {
return fmt.Errorf("error starting algorithm: %v", err)
}
if err := cmd.Wait(); err != nil {
if err := b.cmd.Wait(); err != nil {
return fmt.Errorf("algorithm execution error: %v", err)
}
return nil
}
func (b *binary) Stop() error {
if b.cmd == nil {
return nil
}
if b.cmd.ProcessState != nil && b.cmd.ProcessState.Exited() {
return nil
}
if err := b.cmd.Process.Kill(); err != nil {
return fmt.Errorf("error stopping algorithm: %v", err)
}
return nil
}
+2 -2
View File
@@ -18,7 +18,7 @@ func TestNewAlgorithm(t *testing.T) {
algoFile := "/path/to/algo"
args := []string{"arg1", "arg2"}
algo := NewAlgorithm(logger, eventsSvc, algoFile, args)
algo := NewAlgorithm(logger, eventsSvc, algoFile, args, "")
b, ok := algo.(*binary)
if !ok {
@@ -74,7 +74,7 @@ func TestBinaryRun(t *testing.T) {
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
eventsSvc := new(mocks.Service)
b := NewAlgorithm(logger, eventsSvc, tt.algoFile, tt.args).(*binary)
b := NewAlgorithm(logger, eventsSvc, tt.algoFile, tt.args, "").(*binary)
var stdout, stderr bytes.Buffer
b.stdout = &stdout
+8 -4
View File
@@ -35,11 +35,11 @@ type docker struct {
stdout io.Writer
}
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string) algorithm.Algorithm {
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile, cmpID string) algorithm.Algorithm {
d := &docker{
algoFile: algoFile,
logger: logger,
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc},
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc, CmpID: cmpID},
stdout: &logging.Stdout{Logger: logger},
}
@@ -47,8 +47,6 @@ func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string
}
func (d *docker) Run() error {
ctx := context.Background()
// Create a new Docker client.
cli, err := client.NewClientWithOpts(client.FromEnv, client.WithAPIVersionNegotiation())
if err != nil {
@@ -62,6 +60,7 @@ func (d *docker) Run() error {
}
defer imageFile.Close()
ctx := context.Background()
// Load the Docker image from the tar file.
resp, err := cli.ImageLoad(ctx, imageFile, true)
if err != nil {
@@ -176,3 +175,8 @@ func writeToOut(readCloser io.ReadCloser, ioWriter io.Writer) error {
return nil
}
func (d *docker) Stop() error {
// To be supported later.
return nil
}
+1 -1
View File
@@ -18,7 +18,7 @@ func TestNewAlgorithm(t *testing.T) {
eventsSvc := new(mocks.Service)
algoFile := "/path/to/algo.tar"
algo := NewAlgorithm(logger, eventsSvc, algoFile)
algo := NewAlgorithm(logger, eventsSvc, algoFile, "")
d, ok := algo.(*docker)
assert.True(t, ok, "NewAlgorithm should return a *docker")
+2 -3
View File
@@ -50,6 +50,7 @@ func (s *Stdout) Write(p []byte) (n int, err error) {
type Stderr struct {
Logger *slog.Logger
EventSvc events.Service
CmpID string
}
// Write implements io.Writer.
@@ -70,9 +71,7 @@ func (s *Stderr) Write(p []byte) (n int, err error) {
s.Logger.Error(string(buf[:n]))
}
if err := s.EventSvc.SendEvent(algorithmRun, warningStatus, json.RawMessage{}); err != nil {
return len(p), err
}
s.EventSvc.SendEvent(s.CmpID, algorithmRun, warningStatus, json.RawMessage{})
return len(p), nil
}
+1 -1
View File
@@ -73,7 +73,7 @@ func TestStderrWrite(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockEventService := mocks.NewService(t)
mockEventService.On("SendEvent", "AlgorithmRun", manager.Warning.String(), mock.Anything).Return(nil)
mockEventService.On("SendEvent", mock.Anything, "AlgorithmRun", manager.Warning.String(), mock.Anything).Return(nil)
stderr := &Stderr{Logger: mglog.NewMock(), EventSvc: mockEventService}
n, err := stderr.Write([]byte(tt.input))
+24 -7
View File
@@ -39,12 +39,13 @@ type python struct {
runtime string
requirementsFile string
args []string
cmd *exec.Cmd
}
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, runtime, requirementsFile, algoFile string, args []string) algorithm.Algorithm {
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, runtime, requirementsFile, algoFile string, args []string, cmpID string) algorithm.Algorithm {
p := &python{
algoFile: algoFile,
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc},
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc, CmpID: cmpID},
stdout: &logging.Stdout{Logger: logger},
requirementsFile: requirementsFile,
args: args,
@@ -85,15 +86,15 @@ func (p *python) Run() error {
}
args := append([]string{p.algoFile}, p.args...)
cmd := exec.Command(pythonPath, args...)
cmd.Stderr = p.stderr
cmd.Stdout = p.stdout
p.cmd = exec.Command(pythonPath, args...)
p.cmd.Stderr = p.stderr
p.cmd.Stdout = p.stdout
if err := cmd.Start(); err != nil {
if err := p.cmd.Start(); err != nil {
return fmt.Errorf("error starting algorithm: %v", err)
}
if err := cmd.Wait(); err != nil {
if err := p.cmd.Wait(); err != nil {
return fmt.Errorf("algorithm execution error: %v", err)
}
@@ -103,3 +104,19 @@ func (p *python) Run() error {
return nil
}
func (p *python) Stop() error {
if p.cmd == nil {
return nil
}
if p.cmd.ProcessState != nil && p.cmd.ProcessState.Exited() {
return nil
}
if err := p.cmd.Process.Kill(); err != nil {
return fmt.Errorf("error stopping algorithm: %v", err)
}
return nil
}
+1 -1
View File
@@ -50,7 +50,7 @@ func TestNewAlgorithm(t *testing.T) {
algoFile := "algorithm.py"
args := []string{"--arg1", "value1"}
algo := NewAlgorithm(logger, eventsSvc, runtime, requirementsFile, algoFile, args)
algo := NewAlgorithm(logger, eventsSvc, runtime, requirementsFile, algoFile, args, "")
p, ok := algo.(*python)
if !ok {
+24 -7
View File
@@ -24,12 +24,13 @@ type wasm struct {
stderr io.Writer
stdout io.Writer
args []string
cmd *exec.Cmd
}
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string, args []string) algorithm.Algorithm {
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, args []string, algoFile, cmpID string) algorithm.Algorithm {
return &wasm{
algoFile: algoFile,
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc},
stderr: &logging.Stderr{Logger: logger, EventSvc: eventsSvc, CmpID: cmpID},
stdout: &logging.Stdout{Logger: logger},
args: args,
}
@@ -38,17 +39,33 @@ func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string
func (w *wasm) Run() error {
args := append(mapDirOption, w.algoFile)
args = append(args, w.args...)
cmd := exec.Command(wasmRuntime, args...)
cmd.Stderr = w.stderr
cmd.Stdout = w.stdout
w.cmd = exec.Command(wasmRuntime, args...)
w.cmd.Stderr = w.stderr
w.cmd.Stdout = w.stdout
if err := cmd.Start(); err != nil {
if err := w.cmd.Start(); err != nil {
return fmt.Errorf("error starting algorithm: %v", err)
}
if err := cmd.Wait(); err != nil {
if err := w.cmd.Wait(); err != nil {
return fmt.Errorf("algorithm execution error: %v", err)
}
return nil
}
func (w *wasm) Stop() error {
if w.cmd == nil {
return nil
}
if w.cmd.ProcessState != nil && w.cmd.ProcessState.Exited() {
return nil
}
if err := w.cmd.Process.Kill(); err != nil {
return fmt.Errorf("error stopping algorithm: %v", err)
}
return nil
}
+2 -2
View File
@@ -18,7 +18,7 @@ func TestNewAlgorithm(t *testing.T) {
algoFile := "test.wasm"
args := []string{"arg1", "arg2"}
algo := NewAlgorithm(logger, eventsSvc, algoFile, args)
algo := NewAlgorithm(logger, eventsSvc, args, algoFile, "")
w, ok := algo.(*wasm)
if !ok {
@@ -54,7 +54,7 @@ func TestRunError(t *testing.T) {
algoFile := "test.wasm"
args := []string{"arg1", "arg2"}
w := NewAlgorithm(logger, eventsSvc, algoFile, args).(*wasm)
w := NewAlgorithm(logger, eventsSvc, args, algoFile, "").(*wasm)
err := w.Run()
if err == nil {
+28
View File
@@ -27,6 +27,34 @@ func LoggingMiddleware(svc agent.Service, logger *slog.Logger) agent.Service {
return &loggingMiddleware{logger, svc}
}
// InitComputation implements agent.Service.
func (lm *loggingMiddleware) InitComputation(ctx context.Context, cmp agent.Computation) (err error) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method InitComputation for computation id %s took %s to complete", cmp.ID, time.Since(begin))
if err != nil {
lm.logger.WithGroup(cmp.ID).Warn(fmt.Sprintf("%s with error: %s", message, err))
return
}
lm.logger.WithGroup(cmp.ID).Info(fmt.Sprintf("%s without errors", message))
}(time.Now())
return lm.svc.InitComputation(ctx, cmp)
}
// StopComputation implements agent.Service.
func (lm *loggingMiddleware) StopComputation(ctx context.Context) (err error) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method StopComputation took %s to complete", time.Since(begin))
if err != nil {
lm.logger.Warn(fmt.Sprintf("%s with error: %s", message, err))
return
}
lm.logger.Info(fmt.Sprintf("%s without errors", message))
}(time.Now())
return lm.svc.StopComputation(ctx)
}
func (lm *loggingMiddleware) Algo(ctx context.Context, algorithm agent.Algorithm) (err error) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method Algo took %s to complete", time.Since(begin))
+20
View File
@@ -32,6 +32,26 @@ func MetricsMiddleware(svc agent.Service, counter metrics.Counter, latency metri
}
}
// InitComputation implements agent.Service.
func (ms *metricsMiddleware) InitComputation(ctx context.Context, cmp agent.Computation) error {
defer func(begin time.Time) {
ms.counter.With("method", "init_computation").Add(1)
ms.latency.With("method", "init_computation").Observe(time.Since(begin).Seconds())
}(time.Now())
return ms.svc.InitComputation(ctx, cmp)
}
// StopComputation implements agent.Service.
func (ms *metricsMiddleware) StopComputation(ctx context.Context) error {
defer func(begin time.Time) {
ms.counter.With("method", "stop_computation").Add(1)
ms.latency.With("method", "stop_computation").Observe(time.Since(begin).Seconds())
}(time.Now())
return ms.svc.StopComputation(ctx)
}
func (ms *metricsMiddleware) Algo(ctx context.Context, algorithm agent.Algorithm) error {
defer func(begin time.Time) {
ms.counter.With("method", "algo").Add(1)
-2
View File
@@ -13,7 +13,6 @@ import (
var _ fmt.Stringer = (*Datasets)(nil)
type AgentConfig struct {
LogLevel string `json:"log_level,omitempty"`
Host string `json:"host,omitempty"`
Port string `json:"port,omitempty"`
CertFile string `json:"cert_file,omitempty"`
@@ -30,7 +29,6 @@ type Computation struct {
Datasets Datasets `json:"datasets,omitempty"`
Algorithm Algorithm `json:"algorithm,omitempty"`
ResultConsumers []ResultConsumer `json:"result_consumers,omitempty"`
AgentConfig AgentConfig `json:"agent_config,omitempty"`
}
type ResultConsumer struct {
-1
View File
@@ -106,7 +106,6 @@ func TestDecompressToContext(t *testing.T) {
func TestAgentConfigJSON(t *testing.T) {
config := AgentConfig{
LogLevel: "info",
Host: "localhost",
Port: "8080",
CertFile: "cert.pem",
+261
View File
@@ -0,0 +1,261 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package grpc
import (
"context"
"log/slog"
"sync"
"time"
"github.com/absmach/magistrala/pkg/errors"
"github.com/ultravioletrs/cocos/agent"
"github.com/ultravioletrs/cocos/agent/cvms"
"github.com/ultravioletrs/cocos/agent/cvms/server"
"golang.org/x/sync/errgroup"
"google.golang.org/protobuf/proto"
)
var (
errCorruptedManifest = errors.New("received manifest may be corrupted")
errUnknonwMessageType = errors.New("unknown message type")
sendTimeout = 5 * time.Second
)
type CVMSClient struct {
mu sync.Mutex
stream cvms.Service_ProcessClient
svc agent.Service
messageQueue chan *cvms.ClientStreamMessage
logger *slog.Logger
runReqManager *runRequestManager
sp server.AgentServer
}
// NewClient returns new gRPC client instance.
func NewClient(stream cvms.Service_ProcessClient, svc agent.Service, messageQueue chan *cvms.ClientStreamMessage, logger *slog.Logger, sp server.AgentServer) CVMSClient {
return CVMSClient{
stream: stream,
svc: svc,
messageQueue: messageQueue,
logger: logger,
runReqManager: newRunRequestManager(),
sp: sp,
}
}
func (client *CVMSClient) Process(ctx context.Context, cancel context.CancelFunc) error {
eg, ctx := errgroup.WithContext(ctx)
eg.Go(func() error {
return client.handleIncomingMessages(ctx)
})
eg.Go(func() error {
return client.handleOutgoingMessages(ctx)
})
return eg.Wait()
}
func (client *CVMSClient) handleIncomingMessages(ctx context.Context) error {
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
req, err := client.stream.Recv()
if err != nil {
return err
}
if err := client.processIncomingMessage(ctx, req); err != nil {
return err
}
}
}
}
func (client *CVMSClient) processIncomingMessage(ctx context.Context, req *cvms.ServerStreamMessage) error {
switch mes := req.Message.(type) {
case *cvms.ServerStreamMessage_RunReqChunks:
return client.handleRunReqChunks(ctx, mes)
case *cvms.ServerStreamMessage_StopComputation:
go client.handleStopComputation(ctx, mes)
default:
return errUnknonwMessageType
}
return nil
}
func (client *CVMSClient) handleRunReqChunks(ctx context.Context, msg *cvms.ServerStreamMessage_RunReqChunks) error {
buffer, complete := client.runReqManager.addChunk(msg.RunReqChunks.Id, msg.RunReqChunks.Data, msg.RunReqChunks.IsLast)
if complete {
var runReq cvms.ComputationRunReq
if err := proto.Unmarshal(buffer, &runReq); err != nil {
return errors.Wrap(err, errCorruptedManifest)
}
go client.executeRun(ctx, &runReq)
}
return nil
}
func (client *CVMSClient) executeRun(ctx context.Context, runReq *cvms.ComputationRunReq) {
ac := agent.Computation{
ID: runReq.Id,
Name: runReq.Name,
Description: runReq.Description,
}
if runReq.Algorithm != nil {
ac.Algorithm = agent.Algorithm{
Hash: [32]byte(runReq.Algorithm.Hash),
UserKey: runReq.Algorithm.UserKey,
}
}
for _, ds := range runReq.Datasets {
ac.Datasets = append(ac.Datasets, agent.Dataset{
Hash: [32]byte(ds.Hash),
UserKey: ds.UserKey,
})
}
for _, rc := range runReq.ResultConsumers {
ac.ResultConsumers = append(ac.ResultConsumers, agent.ResultConsumer{
UserKey: rc.UserKey,
})
}
if err := client.svc.InitComputation(ctx, ac); err != nil {
client.logger.Warn(err.Error())
return
}
client.mu.Lock()
defer client.mu.Unlock()
if runReq.AgentConfig == nil {
runReq.AgentConfig = &cvms.AgentConfig{}
}
runRes := &cvms.ClientStreamMessage_RunRes{
RunRes: &cvms.RunResponse{
ComputationId: runReq.Id,
},
}
err := client.sp.Start(ctx, agent.AgentConfig{
Port: runReq.AgentConfig.Port,
Host: runReq.AgentConfig.Host,
CertFile: runReq.AgentConfig.CertFile,
KeyFile: runReq.AgentConfig.KeyFile,
ServerCAFile: runReq.AgentConfig.ServerCaFile,
ClientCAFile: runReq.AgentConfig.ClientCaFile,
AttestedTls: runReq.AgentConfig.AttestedTls,
}, ac)
if err != nil {
client.logger.Warn(err.Error())
runRes.RunRes.Error = err.Error()
}
client.sendMessage(&cvms.ClientStreamMessage{Message: runRes})
}
func (client *CVMSClient) handleStopComputation(ctx context.Context, mes *cvms.ServerStreamMessage_StopComputation) {
msg := &cvms.ClientStreamMessage_StopComputationRes{
StopComputationRes: &cvms.StopComputationResponse{
ComputationId: mes.StopComputation.ComputationId,
},
}
if err := client.svc.StopComputation(ctx); err != nil {
msg.StopComputationRes.Message = err.Error()
}
client.mu.Lock()
defer client.mu.Unlock()
if err := client.sp.Stop(); err != nil {
msg.StopComputationRes.Message = err.Error()
}
client.sendMessage(&cvms.ClientStreamMessage{Message: msg})
}
func (client *CVMSClient) handleOutgoingMessages(ctx context.Context) error {
for {
select {
case <-ctx.Done():
return ctx.Err()
case mes := <-client.messageQueue:
if err := client.stream.Send(mes); err != nil {
return err
}
}
}
}
func (client *CVMSClient) sendMessage(mes *cvms.ClientStreamMessage) {
ctx, cancel := context.WithTimeout(context.Background(), sendTimeout)
defer cancel()
select {
case client.messageQueue <- mes:
case <-ctx.Done():
client.logger.Warn("Failed to send message: timeout exceeded")
}
}
type runRequestManager struct {
requests map[string]*runRequest
mu sync.Mutex
}
type runRequest struct {
buffer []byte
lastChunk time.Time
timer *time.Timer
}
func newRunRequestManager() *runRequestManager {
return &runRequestManager{
requests: make(map[string]*runRequest),
}
}
func (m *runRequestManager) addChunk(id string, chunk []byte, isLast bool) ([]byte, bool) {
m.mu.Lock()
defer m.mu.Unlock()
req, exists := m.requests[id]
if !exists {
req = &runRequest{
buffer: make([]byte, 0),
lastChunk: time.Now(),
timer: time.AfterFunc(runReqTimeout, func() { m.timeoutRequest(id) }),
}
m.requests[id] = req
}
req.buffer = append(req.buffer, chunk...)
req.lastChunk = time.Now()
req.timer.Reset(runReqTimeout)
if isLast {
delete(m.requests, id)
req.timer.Stop()
return req.buffer, true
}
return nil, false
}
func (m *runRequestManager) timeoutRequest(id string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.requests, id)
// Log timeout or handle it as needed
}
+202
View File
@@ -0,0 +1,202 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package grpc
import (
"context"
"testing"
"time"
mglog "github.com/absmach/magistrala/logger"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/ultravioletrs/cocos/agent/cvms"
servermocks "github.com/ultravioletrs/cocos/agent/cvms/server/mocks"
"github.com/ultravioletrs/cocos/agent/mocks"
"google.golang.org/grpc"
"google.golang.org/protobuf/proto"
)
type mockStream struct {
mock.Mock
grpc.ClientStream
}
func (m *mockStream) Recv() (*cvms.ServerStreamMessage, error) {
args := m.Called()
return args.Get(0).(*cvms.ServerStreamMessage), args.Error(1)
}
func (m *mockStream) Send(msg *cvms.ClientStreamMessage) error {
args := m.Called(msg)
return args.Error(0)
}
func TestManagerClient_Process1(t *testing.T) {
tests := []struct {
name string
setupMocks func(mockStream *mockStream, mockSvc *mocks.Service, mockServerSvc *servermocks.AgentServerProvider)
expectError bool
errorMsg string
}{
{
name: "Stop computation",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service, mockServerSvc *servermocks.AgentServerProvider) {
mockStream.On("Recv").Return(&cvms.ServerStreamMessage{
Message: &cvms.ServerStreamMessage_StopComputation{
StopComputation: &cvms.StopComputation{},
},
}, nil)
mockStream.On("Send", mock.Anything).Return(nil)
mockSvc.On("StopComputation", mock.Anything).Return(nil)
mockServerSvc.On("Stop").Return(nil)
},
expectError: true,
errorMsg: "context deadline exceeded",
},
{
name: "Run request chunks",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service, mockServerSvc *servermocks.AgentServerProvider) {
mockStream.On("Recv").Return(&cvms.ServerStreamMessage{
Message: &cvms.ServerStreamMessage_RunReqChunks{
RunReqChunks: &cvms.RunReqChunks{},
},
}, nil)
mockStream.On("Send", mock.Anything).Return(nil).Once()
mockSvc.On("Run", mock.Anything, mock.Anything).Return("", assert.AnError).Once()
},
expectError: true,
},
{
name: "Receive error",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service, mockServerSvc *servermocks.AgentServerProvider) {
mockStream.On("Recv").Return(&cvms.ServerStreamMessage{}, assert.AnError)
},
expectError: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
mockServerSvc := new(servermocks.AgentServerProvider)
messageQueue := make(chan *cvms.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger, mockServerSvc)
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
tc.setupMocks(mockStream, mockSvc, mockServerSvc)
err := client.Process(ctx, cancel)
if tc.expectError {
assert.Error(t, err)
if tc.errorMsg != "" {
assert.Contains(t, err.Error(), tc.errorMsg)
}
} else {
assert.NoError(t, err)
}
})
}
}
func TestManagerClient_handleRunReqChunks(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
mockServerSvc := new(servermocks.AgentServerProvider)
messageQueue := make(chan *cvms.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger, mockServerSvc)
runReq := &cvms.ComputationRunReq{
Id: "test-id",
}
runReqBytes, _ := proto.Marshal(runReq)
chunk1 := &cvms.ServerStreamMessage_RunReqChunks{
RunReqChunks: &cvms.RunReqChunks{
Id: "chunk-1",
Data: runReqBytes[:len(runReqBytes)/2],
IsLast: false,
},
}
chunk2 := &cvms.ServerStreamMessage_RunReqChunks{
RunReqChunks: &cvms.RunReqChunks{
Id: "chunk-1",
Data: runReqBytes[len(runReqBytes)/2:],
IsLast: true,
},
}
mockSvc.On("InitComputation", mock.Anything, mock.Anything).Return(nil)
mockServerSvc.On("Start", mock.Anything, mock.Anything, mock.Anything).Return(nil)
err := client.handleRunReqChunks(context.Background(), chunk1)
assert.NoError(t, err)
err = client.handleRunReqChunks(context.Background(), chunk2)
assert.NoError(t, err)
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
runRes, ok := msg.Message.(*cvms.ClientStreamMessage_RunRes)
assert.True(t, ok)
assert.Equal(t, "test-id", runRes.RunRes.ComputationId)
}
func TestManagerClient_handleStopComputation(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
mockServerSvc := new(servermocks.AgentServerProvider)
messageQueue := make(chan *cvms.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger, mockServerSvc)
stopReq := &cvms.ServerStreamMessage_StopComputation{
StopComputation: &cvms.StopComputation{
ComputationId: "test-comp-id",
},
}
mockSvc.On("StopComputation", mock.Anything).Return(nil)
mockServerSvc.On("Stop").Return(nil)
client.handleStopComputation(context.Background(), stopReq)
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
stopRes, ok := msg.Message.(*cvms.ClientStreamMessage_StopComputationRes)
assert.True(t, ok)
assert.Equal(t, "test-comp-id", stopRes.StopComputationRes.ComputationId)
assert.Empty(t, stopRes.StopComputationRes.Message)
}
func TestManagerClient_timeoutRequest(t *testing.T) {
rm := newRunRequestManager()
rm.requests["test-id"] = &runRequest{
timer: time.NewTimer(100 * time.Millisecond),
buffer: []byte("test-data"),
lastChunk: time.Now(),
}
rm.timeoutRequest("test-id")
assert.Len(t, rm.requests, 0)
}
+5
View File
@@ -0,0 +1,5 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
// Package grpc contains implementation of kit service gRPC API.
package grpc
+133
View File
@@ -0,0 +1,133 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package grpc
import (
"bytes"
"context"
"errors"
"io"
"time"
"github.com/ultravioletrs/cocos/agent/cvms"
"golang.org/x/sync/errgroup"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
"google.golang.org/protobuf/proto"
)
var (
_ cvms.ServiceServer = (*grpcServer)(nil)
ErrUnexpectedMsg = errors.New("unknown message type")
)
const (
bufferSize = 1024 * 1024 // 1 MB
runReqTimeout = 30 * time.Second
)
type SendFunc func(*cvms.ServerStreamMessage) error
type grpcServer struct {
cvms.UnimplementedServiceServer
incoming chan *cvms.ClientStreamMessage
svc Service
}
type Service interface {
Run(ctx context.Context, ipAddress string, sendMessage SendFunc, authInfo credentials.AuthInfo)
}
// NewServer returns new AuthServiceServer instance.
func NewServer(incoming chan *cvms.ClientStreamMessage, svc Service) cvms.ServiceServer {
return &grpcServer{
incoming: incoming,
svc: svc,
}
}
func (s *grpcServer) Process(stream cvms.Service_ProcessServer) error {
client, ok := peer.FromContext(stream.Context())
if !ok {
return errors.New("failed to get peer info")
}
eg, ctx := errgroup.WithContext(stream.Context())
eg.Go(func() error {
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
req, err := stream.Recv()
if err != nil {
return err
}
s.incoming <- req
}
}
})
eg.Go(func() error {
sendMessage := func(msg *cvms.ServerStreamMessage) error {
select {
case <-ctx.Done():
return ctx.Err()
default:
switch m := msg.Message.(type) {
case *cvms.ServerStreamMessage_RunReq:
return s.sendRunReqInChunks(stream, m.RunReq)
default:
return stream.Send(msg)
}
}
}
s.svc.Run(ctx, client.Addr.String(), sendMessage, client.AuthInfo)
return nil
})
return eg.Wait()
}
func (s *grpcServer) sendRunReqInChunks(stream cvms.Service_ProcessServer, runReq *cvms.ComputationRunReq) error {
data, err := proto.Marshal(runReq)
if err != nil {
return err
}
dataBuffer := bytes.NewBuffer(data)
buf := make([]byte, bufferSize)
for {
n, err := dataBuffer.Read(buf)
isLast := false
if err == io.EOF {
isLast = true
} else if err != nil {
return err
}
chunk := &cvms.ServerStreamMessage{
Message: &cvms.ServerStreamMessage_RunReqChunks{
RunReqChunks: &cvms.RunReqChunks{
Id: runReq.Id,
Data: buf[:n],
IsLast: isLast,
},
},
}
if err := stream.Send(chunk); err != nil {
return err
}
if isLast {
break
}
}
return nil
}
@@ -10,24 +10,24 @@ import (
"github.com/absmach/magistrala/pkg/errors"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/ultravioletrs/cocos/manager"
"github.com/ultravioletrs/cocos/agent/cvms"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
)
type mockServerStream struct {
mock.Mock
manager.ManagerService_ProcessServer
cvms.Service_ProcessServer
}
func (m *mockServerStream) Send(msg *manager.ServerStreamMessage) error {
func (m *mockServerStream) Send(msg *cvms.ServerStreamMessage) error {
args := m.Called(msg)
return args.Error(0)
}
func (m *mockServerStream) Recv() (*manager.ClientStreamMessage, error) {
func (m *mockServerStream) Recv() (*cvms.ClientStreamMessage, error) {
args := m.Called()
return args.Get(0).(*manager.ClientStreamMessage), args.Error(1)
return args.Get(0).(*cvms.ClientStreamMessage), args.Error(1)
}
func (m *mockServerStream) Context() context.Context {
@@ -44,7 +44,7 @@ func (m *mockService) Run(ctx context.Context, ipAddress string, sendMessage Sen
}
func TestNewServer(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage)
incoming := make(chan *cvms.ClientStreamMessage)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc)
@@ -56,19 +56,19 @@ func TestNewServer(t *testing.T) {
func TestGrpcServer_Process(t *testing.T) {
tests := []struct {
name string
recvReturn *manager.ClientStreamMessage
recvReturn *cvms.ClientStreamMessage
recvError error
expectedError string
}{
{
name: "Process with context deadline exceeded",
recvReturn: &manager.ClientStreamMessage{},
recvReturn: &cvms.ClientStreamMessage{},
recvError: nil,
expectedError: "context deadline exceeded",
},
{
name: "Process with Recv error",
recvReturn: &manager.ClientStreamMessage{},
recvReturn: &cvms.ClientStreamMessage{},
recvError: errors.New("recv error"),
expectedError: "recv error",
},
@@ -76,7 +76,7 @@ func TestGrpcServer_Process(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage, 1)
incoming := make(chan *cvms.ClientStreamMessage, 1)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc).(*grpcServer)
@@ -111,13 +111,13 @@ func TestGrpcServer_Process(t *testing.T) {
}
func TestGrpcServer_sendRunReqInChunks(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage)
incoming := make(chan *cvms.ClientStreamMessage)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc).(*grpcServer)
mockStream := new(mockServerStream)
runReq := &manager.ComputationRunReq{
runReq := &cvms.ComputationRunReq{
Id: "test-id",
}
@@ -125,10 +125,10 @@ func TestGrpcServer_sendRunReqInChunks(t *testing.T) {
for i := range largePayload {
largePayload[i] = byte(i % 256)
}
runReq.Algorithm = &manager.Algorithm{}
runReq.Algorithm = &cvms.Algorithm{}
runReq.Algorithm.UserKey = largePayload
mockStream.On("Send", mock.AnythingOfType("*manager.ServerStreamMessage")).Return(nil).Times(4)
mockStream.On("Send", mock.AnythingOfType("*cvms.ServerStreamMessage")).Return(nil).Times(4)
err := server.sendRunReqInChunks(mockStream, runReq)
@@ -139,7 +139,7 @@ func TestGrpcServer_sendRunReqInChunks(t *testing.T) {
assert.Equal(t, 4, len(calls))
for i, call := range calls {
msg := call.Arguments[0].(*manager.ServerStreamMessage)
msg := call.Arguments[0].(*cvms.ServerStreamMessage)
chunk := msg.GetRunReqChunks()
assert.NotNil(t, chunk)
@@ -174,9 +174,9 @@ func TestGrpcServer_ProcessWithMockService(t *testing.T) {
mockSvc.On("Run", mock.Anything, "test", mock.Anything, mock.AnythingOfType("mockAuthInfo")).
Run(func(args mock.Arguments) {
sendFunc := args.Get(2).(SendFunc)
runReq := &manager.ComputationRunReq{Id: "test-run-id"}
err := sendFunc(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_RunReq{
runReq := &cvms.ComputationRunReq{Id: "test-run-id"}
err := sendFunc(&cvms.ServerStreamMessage{
Message: &cvms.ServerStreamMessage_RunReq{
RunReq: runReq,
},
})
@@ -184,34 +184,17 @@ func TestGrpcServer_ProcessWithMockService(t *testing.T) {
}).
Return()
mockStream.On("Send", mock.MatchedBy(func(msg *manager.ServerStreamMessage) bool {
mockStream.On("Send", mock.MatchedBy(func(msg *cvms.ServerStreamMessage) bool {
chunks := msg.GetRunReqChunks()
return chunks != nil && chunks.Id == "test-run-id"
})).Return(nil)
},
},
{
name: "Terminate Request Test",
setupMockFn: func(mockSvc *mockService, mockStream *mockServerStream) {
mockSvc.On("Run", mock.Anything, "test", mock.Anything, mock.AnythingOfType("mockAuthInfo")).
Run(func(args mock.Arguments) {
sendFunc := args.Get(2).(SendFunc)
err := sendFunc(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_TerminateReq{
TerminateReq: &manager.Terminate{},
},
})
assert.NoError(t, err)
}).Return()
mockStream.On("Send", mock.AnythingOfType("*manager.ServerStreamMessage")).Return(nil)
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage, 10)
incoming := make(chan *cvms.ClientStreamMessage, 10)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc).(*grpcServer)
@@ -231,7 +214,7 @@ func TestGrpcServer_ProcessWithMockService(t *testing.T) {
})
mockStream.On("Context").Return(peerCtx)
mockStream.On("Recv").Return(&manager.ClientStreamMessage{}, nil).Maybe()
mockStream.On("Recv").Return(&cvms.ClientStreamMessage{}, nil).Maybe()
tt.setupMockFn(mockSvc, mockStream)
@@ -251,18 +234,18 @@ func TestGrpcServer_ProcessWithMockService(t *testing.T) {
}
func TestGrpcServer_sendRunReqInChunksError(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage)
incoming := make(chan *cvms.ClientStreamMessage)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc).(*grpcServer)
mockStream := new(mockServerStream)
runReq := &manager.ComputationRunReq{
runReq := &cvms.ComputationRunReq{
Id: "test-id",
}
// Simulate an error when sending
mockStream.On("Send", mock.AnythingOfType("*manager.ServerStreamMessage")).Return(errors.New("send error")).Once()
mockStream.On("Send", mock.AnythingOfType("*cvms.ServerStreamMessage")).Return(errors.New("send error")).Once()
err := server.sendRunReqInChunks(mockStream, runReq)
@@ -272,7 +255,7 @@ func TestGrpcServer_sendRunReqInChunksError(t *testing.T) {
}
func TestGrpcServer_ProcessMissingPeerInfo(t *testing.T) {
incoming := make(chan *manager.ClientStreamMessage)
incoming := make(chan *cvms.ClientStreamMessage)
mockSvc := new(mockService)
server := NewServer(incoming, mockSvc).(*grpcServer)
File diff suppressed because it is too large Load Diff
+103
View File
@@ -0,0 +1,103 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
syntax = "proto3";
import "google/protobuf/timestamp.proto";
package cvms;
option go_package = "./cvms";
service Service {
rpc Process(stream ClientStreamMessage) returns (stream ServerStreamMessage) {}
}
message StopComputation {
string computation_id = 1;
}
message StopComputationResponse {
string computation_id = 1;
string message = 2;
}
message RunResponse{
string computation_id = 1;
string error = 2;
}
message AgentEvent {
string event_type = 1;
google.protobuf.Timestamp timestamp = 2;
string computation_id = 3;
bytes details = 4;
string originator = 5;
string status = 6;
}
message AgentLog {
string message = 1;
string computation_id = 2;
string level = 3;
google.protobuf.Timestamp timestamp = 4;
}
message ClientStreamMessage {
oneof message {
AgentLog agent_log = 1;
AgentEvent agent_event = 2;
RunResponse run_res = 3;
StopComputationResponse stopComputationRes = 4;
}
}
message ServerStreamMessage {
oneof message {
RunReqChunks runReqChunks = 1;
ComputationRunReq runReq = 2;
StopComputation stopComputation = 3;
}
}
message RunReqChunks {
bytes data = 1;
string id = 2;
bool is_last = 3;
}
message ComputationRunReq {
string id = 1;
string name = 2;
string description = 3;
repeated Dataset datasets = 4;
Algorithm algorithm = 5;
repeated ResultConsumer result_consumers = 6;
AgentConfig agent_config = 7;
}
message ResultConsumer {
bytes userKey = 1;
}
message Dataset {
bytes hash = 1; // should be sha3.Sum256, 32 byte length.
bytes userKey = 2;
string filename = 3;
}
message Algorithm {
bytes hash = 1; // should be sha3.Sum256, 32 byte length.
bytes userKey = 2;
}
message AgentConfig {
string port = 1;
string host = 2;
string cert_file = 3;
string key_file = 4;
string client_ca_file = 5;
string server_ca_file = 6;
string log_level = 7;
bool attested_tls = 8;
}
+118
View File
@@ -0,0 +1,118 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
// Code generated by protoc-gen-go-grpc. DO NOT EDIT.
// versions:
// - protoc-gen-go-grpc v1.5.1
// - protoc v5.29.0
// source: agent/cvms/cvms.proto
package cvms
import (
context "context"
grpc "google.golang.org/grpc"
codes "google.golang.org/grpc/codes"
status "google.golang.org/grpc/status"
)
// This is a compile-time assertion to ensure that this generated file
// is compatible with the grpc package it is being compiled against.
// Requires gRPC-Go v1.64.0 or later.
const _ = grpc.SupportPackageIsVersion9
const (
Service_Process_FullMethodName = "/cvms.Service/Process"
)
// ServiceClient is the client API for Service service.
//
// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream.
type ServiceClient interface {
Process(ctx context.Context, opts ...grpc.CallOption) (grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage], error)
}
type serviceClient struct {
cc grpc.ClientConnInterface
}
func NewServiceClient(cc grpc.ClientConnInterface) ServiceClient {
return &serviceClient{cc}
}
func (c *serviceClient) Process(ctx context.Context, opts ...grpc.CallOption) (grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage], error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
stream, err := c.cc.NewStream(ctx, &Service_ServiceDesc.Streams[0], Service_Process_FullMethodName, cOpts...)
if err != nil {
return nil, err
}
x := &grpc.GenericClientStream[ClientStreamMessage, ServerStreamMessage]{ClientStream: stream}
return x, nil
}
// This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name.
type Service_ProcessClient = grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage]
// ServiceServer is the server API for Service service.
// All implementations must embed UnimplementedServiceServer
// for forward compatibility.
type ServiceServer interface {
Process(grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]) error
mustEmbedUnimplementedServiceServer()
}
// UnimplementedServiceServer must be embedded to have
// forward compatible implementations.
//
// NOTE: this should be embedded by value instead of pointer to avoid a nil
// pointer dereference when methods are called.
type UnimplementedServiceServer struct{}
func (UnimplementedServiceServer) Process(grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]) error {
return status.Errorf(codes.Unimplemented, "method Process not implemented")
}
func (UnimplementedServiceServer) mustEmbedUnimplementedServiceServer() {}
func (UnimplementedServiceServer) testEmbeddedByValue() {}
// UnsafeServiceServer may be embedded to opt out of forward compatibility for this service.
// Use of this interface is not recommended, as added methods to ServiceServer will
// result in compilation errors.
type UnsafeServiceServer interface {
mustEmbedUnimplementedServiceServer()
}
func RegisterServiceServer(s grpc.ServiceRegistrar, srv ServiceServer) {
// If the following call pancis, it indicates UnimplementedServiceServer was
// embedded by pointer and is nil. This will cause panics if an
// unimplemented method is ever invoked, so we test this at initialization
// time to prevent it from happening at runtime later due to I/O.
if t, ok := srv.(interface{ testEmbeddedByValue() }); ok {
t.testEmbeddedByValue()
}
s.RegisterService(&Service_ServiceDesc, srv)
}
func _Service_Process_Handler(srv interface{}, stream grpc.ServerStream) error {
return srv.(ServiceServer).Process(&grpc.GenericServerStream[ClientStreamMessage, ServerStreamMessage]{ServerStream: stream})
}
// This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name.
type Service_ProcessServer = grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]
// Service_ServiceDesc is the grpc.ServiceDesc for Service service.
// It's only intended for direct use with grpc.RegisterService,
// and not to be introspected or modified (even as a copy)
var Service_ServiceDesc = grpc.ServiceDesc{
ServiceName: "cvms.Service",
HandlerType: (*ServiceServer)(nil),
Methods: []grpc.MethodDesc{},
Streams: []grpc.StreamDesc{
{
StreamName: "Process",
Handler: _Service_Process_Handler,
ServerStreams: true,
ClientStreams: true,
},
},
Metadata: "agent/cvms/cvms.proto",
}
+89
View File
@@ -0,0 +1,89 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package server
import (
context "context"
"fmt"
"log/slog"
"github.com/ultravioletrs/cocos/agent"
agentgrpc "github.com/ultravioletrs/cocos/agent/api/grpc"
"github.com/ultravioletrs/cocos/agent/auth"
"github.com/ultravioletrs/cocos/internal/server"
grpcserver "github.com/ultravioletrs/cocos/internal/server/grpc"
"github.com/ultravioletrs/cocos/pkg/attestation/quoteprovider"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
)
const (
svcName = "agent"
defSvcGRPCPort = "7002"
)
type AgentServer interface {
Start(ctx context.Context, cfg agent.AgentConfig, cmp agent.Computation) error
Stop() error
}
type agentServer struct {
gs server.Server
logger *slog.Logger
svc agent.Service
}
func NewServer(logger *slog.Logger, svc agent.Service) AgentServer {
return &agentServer{
logger: logger,
svc: svc,
}
}
func (as *agentServer) Start(ctx context.Context, cfg agent.AgentConfig, cmp agent.Computation) error {
if cfg.Port == "" {
cfg.Port = defSvcGRPCPort
}
agentGrpcServerConfig := server.AgentConfig{
ServerConfig: server.ServerConfig{
BaseConfig: server.BaseConfig{
Host: cfg.Host,
Port: cfg.Port,
CertFile: cfg.CertFile,
KeyFile: cfg.KeyFile,
ServerCAFile: cfg.ServerCAFile,
ClientCAFile: cfg.ClientCAFile,
},
},
AttestedTLS: cfg.AttestedTls,
}
registerAgentServiceServer := func(srv *grpc.Server) {
reflection.Register(srv)
agent.RegisterAgentServiceServer(srv, agentgrpc.NewServer(as.svc))
}
authSvc, err := auth.New(cmp)
if err != nil {
as.logger.WithGroup(cmp.ID).Error(fmt.Sprintf("failed to create auth service %s", err.Error()))
return err
}
qp, err := quoteprovider.GetQuoteProvider()
if err != nil {
as.logger.Error(fmt.Sprintf("failed to create quote provider %s", err.Error()))
return err
}
ctx, cancel := context.WithCancel(ctx)
as.gs = grpcserver.New(ctx, cancel, svcName, agentGrpcServerConfig, registerAgentServiceServer, as.logger, qp, authSvc)
return as.gs.Start()
}
func (as *agentServer) Stop() error {
return as.gs.Stop()
}
+134
View File
@@ -0,0 +1,134 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
// Code generated by mockery v2.43.2. DO NOT EDIT.
package mocks
import (
context "context"
agent "github.com/ultravioletrs/cocos/agent"
mock "github.com/stretchr/testify/mock"
)
// AgentServerProvider is an autogenerated mock type for the AgentServerProvider type
type AgentServerProvider struct {
mock.Mock
}
type AgentServerProvider_Expecter struct {
mock *mock.Mock
}
func (_m *AgentServerProvider) EXPECT() *AgentServerProvider_Expecter {
return &AgentServerProvider_Expecter{mock: &_m.Mock}
}
// Start provides a mock function with given fields: ctx, cfg, cmp
func (_m *AgentServerProvider) Start(ctx context.Context, cfg agent.AgentConfig, cmp agent.Computation) error {
ret := _m.Called(ctx, cfg, cmp)
if len(ret) == 0 {
panic("no return value specified for Start")
}
var r0 error
if rf, ok := ret.Get(0).(func(context.Context, agent.AgentConfig, agent.Computation) error); ok {
r0 = rf(ctx, cfg, cmp)
} else {
r0 = ret.Error(0)
}
return r0
}
// AgentServerProvider_Start_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Start'
type AgentServerProvider_Start_Call struct {
*mock.Call
}
// Start is a helper method to define mock.On call
// - ctx context.Context
// - cfg agent.AgentConfig
// - cmp agent.Computation
func (_e *AgentServerProvider_Expecter) Start(ctx interface{}, cfg interface{}, cmp interface{}) *AgentServerProvider_Start_Call {
return &AgentServerProvider_Start_Call{Call: _e.mock.On("Start", ctx, cfg, cmp)}
}
func (_c *AgentServerProvider_Start_Call) Run(run func(ctx context.Context, cfg agent.AgentConfig, cmp agent.Computation)) *AgentServerProvider_Start_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context), args[1].(agent.AgentConfig), args[2].(agent.Computation))
})
return _c
}
func (_c *AgentServerProvider_Start_Call) Return(_a0 error) *AgentServerProvider_Start_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *AgentServerProvider_Start_Call) RunAndReturn(run func(context.Context, agent.AgentConfig, agent.Computation) error) *AgentServerProvider_Start_Call {
_c.Call.Return(run)
return _c
}
// Stop provides a mock function with given fields:
func (_m *AgentServerProvider) Stop() error {
ret := _m.Called()
if len(ret) == 0 {
panic("no return value specified for Stop")
}
var r0 error
if rf, ok := ret.Get(0).(func() error); ok {
r0 = rf()
} else {
r0 = ret.Error(0)
}
return r0
}
// AgentServerProvider_Stop_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Stop'
type AgentServerProvider_Stop_Call struct {
*mock.Call
}
// Stop is a helper method to define mock.On call
func (_e *AgentServerProvider_Expecter) Stop() *AgentServerProvider_Stop_Call {
return &AgentServerProvider_Stop_Call{Call: _e.mock.On("Stop")}
}
func (_c *AgentServerProvider_Stop_Call) Run(run func()) *AgentServerProvider_Stop_Call {
_c.Call.Run(func(args mock.Arguments) {
run()
})
return _c
}
func (_c *AgentServerProvider_Stop_Call) Return(_a0 error) *AgentServerProvider_Stop_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *AgentServerProvider_Stop_Call) RunAndReturn(run func() error) *AgentServerProvider_Stop_Call {
_c.Call.Return(run)
return _c
}
// NewAgentServerProvider creates a new instance of AgentServerProvider. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations.
// The first argument is typically a *testing.T value.
func NewAgentServerProvider(t interface {
mock.TestingT
Cleanup(func())
}) *AgentServerProvider {
mock := &AgentServerProvider{}
mock.Mock.Test(t)
t.Cleanup(func() { mock.AssertExpectations(t) })
return mock
}
+19 -24
View File
@@ -4,43 +4,38 @@ package events
import (
"encoding/json"
"io"
"google.golang.org/protobuf/proto"
"github.com/ultravioletrs/cocos/agent/cvms"
"google.golang.org/protobuf/types/known/timestamppb"
)
type service struct {
service string
computationID string
conn io.Writer
service string
queue chan *cvms.ClientStreamMessage
}
type Service interface {
SendEvent(event, status string, details json.RawMessage) error
SendEvent(cmpID, event, status string, details json.RawMessage)
}
func New(svc, computationID string, conn io.Writer) (Service, error) {
func New(svc string, queue chan *cvms.ClientStreamMessage) (Service, error) {
return &service{
service: svc,
computationID: computationID,
conn: conn,
service: svc,
queue: queue,
}, nil
}
func (s *service) SendEvent(event, status string, details json.RawMessage) error {
body := EventsLogs{Message: &EventsLogs_AgentEvent{AgentEvent: &AgentEvent{
EventType: event,
Timestamp: timestamppb.Now(),
ComputationId: s.computationID,
Originator: s.service,
Status: status,
Details: details,
}}}
protoBody, err := proto.Marshal(&body)
if err != nil {
return err
func (s *service) SendEvent(cmpID, event, status string, details json.RawMessage) {
s.queue <- &cvms.ClientStreamMessage{
Message: &cvms.ClientStreamMessage_AgentEvent{
AgentEvent: &cvms.AgentEvent{
EventType: event,
Timestamp: timestamppb.Now(),
ComputationId: cmpID,
Originator: s.service,
Status: status,
Details: details,
},
},
}
_, err = s.conn.Write(protoBody)
return err
}
+36 -79
View File
@@ -3,8 +3,8 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.34.2
// protoc v5.28.1
// protoc-gen-go v1.36.0
// protoc v5.29.0
// source: agent/events/events.proto
package events
@@ -25,25 +25,22 @@ const (
)
type AgentEvent struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
state protoimpl.MessageState `protogen:"open.v1"`
EventType string `protobuf:"bytes,1,opt,name=event_type,json=eventType,proto3" json:"event_type,omitempty"`
Timestamp *timestamppb.Timestamp `protobuf:"bytes,2,opt,name=timestamp,proto3" json:"timestamp,omitempty"`
ComputationId string `protobuf:"bytes,3,opt,name=computation_id,json=computationId,proto3" json:"computation_id,omitempty"`
Details []byte `protobuf:"bytes,4,opt,name=details,proto3" json:"details,omitempty"`
Originator string `protobuf:"bytes,5,opt,name=originator,proto3" json:"originator,omitempty"`
Status string `protobuf:"bytes,6,opt,name=status,proto3" json:"status,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *AgentEvent) Reset() {
*x = AgentEvent{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_events_events_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_events_events_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AgentEvent) String() string {
@@ -54,7 +51,7 @@ func (*AgentEvent) ProtoMessage() {}
func (x *AgentEvent) ProtoReflect() protoreflect.Message {
mi := &file_agent_events_events_proto_msgTypes[0]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -112,23 +109,20 @@ func (x *AgentEvent) GetStatus() string {
}
type AgentLog struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
state protoimpl.MessageState `protogen:"open.v1"`
Message string `protobuf:"bytes,1,opt,name=message,proto3" json:"message,omitempty"`
ComputationId string `protobuf:"bytes,2,opt,name=computation_id,json=computationId,proto3" json:"computation_id,omitempty"`
Level string `protobuf:"bytes,3,opt,name=level,proto3" json:"level,omitempty"`
Timestamp *timestamppb.Timestamp `protobuf:"bytes,4,opt,name=timestamp,proto3" json:"timestamp,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *AgentLog) Reset() {
*x = AgentLog{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_events_events_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_events_events_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *AgentLog) String() string {
@@ -139,7 +133,7 @@ func (*AgentLog) ProtoMessage() {}
func (x *AgentLog) ProtoReflect() protoreflect.Message {
mi := &file_agent_events_events_proto_msgTypes[1]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -183,24 +177,21 @@ func (x *AgentLog) GetTimestamp() *timestamppb.Timestamp {
}
type EventsLogs struct {
state protoimpl.MessageState
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
// Types that are assignable to Message:
state protoimpl.MessageState `protogen:"open.v1"`
// Types that are valid to be assigned to Message:
//
// *EventsLogs_AgentLog
// *EventsLogs_AgentEvent
Message isEventsLogs_Message `protobuf_oneof:"message"`
Message isEventsLogs_Message `protobuf_oneof:"message"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *EventsLogs) Reset() {
*x = EventsLogs{}
if protoimpl.UnsafeEnabled {
mi := &file_agent_events_events_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
mi := &file_agent_events_events_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *EventsLogs) String() string {
@@ -211,7 +202,7 @@ func (*EventsLogs) ProtoMessage() {}
func (x *EventsLogs) ProtoReflect() protoreflect.Message {
mi := &file_agent_events_events_proto_msgTypes[2]
if protoimpl.UnsafeEnabled && x != nil {
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
@@ -226,23 +217,27 @@ func (*EventsLogs) Descriptor() ([]byte, []int) {
return file_agent_events_events_proto_rawDescGZIP(), []int{2}
}
func (m *EventsLogs) GetMessage() isEventsLogs_Message {
if m != nil {
return m.Message
func (x *EventsLogs) GetMessage() isEventsLogs_Message {
if x != nil {
return x.Message
}
return nil
}
func (x *EventsLogs) GetAgentLog() *AgentLog {
if x, ok := x.GetMessage().(*EventsLogs_AgentLog); ok {
return x.AgentLog
if x != nil {
if x, ok := x.Message.(*EventsLogs_AgentLog); ok {
return x.AgentLog
}
}
return nil
}
func (x *EventsLogs) GetAgentEvent() *AgentEvent {
if x, ok := x.GetMessage().(*EventsLogs_AgentEvent); ok {
return x.AgentEvent
if x != nil {
if x, ok := x.Message.(*EventsLogs_AgentEvent); ok {
return x.AgentEvent
}
}
return nil
}
@@ -342,44 +337,6 @@ func file_agent_events_events_proto_init() {
if File_agent_events_events_proto != nil {
return
}
if !protoimpl.UnsafeEnabled {
file_agent_events_events_proto_msgTypes[0].Exporter = func(v any, i int) any {
switch v := v.(*AgentEvent); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_events_events_proto_msgTypes[1].Exporter = func(v any, i int) any {
switch v := v.(*AgentLog); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
file_agent_events_events_proto_msgTypes[2].Exporter = func(v any, i int) any {
switch v := v.(*EventsLogs); i {
case 0:
return &v.state
case 1:
return &v.sizeCache
case 2:
return &v.unknownFields
default:
return nil
}
}
}
file_agent_events_events_proto_msgTypes[2].OneofWrappers = []any{
(*EventsLogs_AgentLog)(nil),
(*EventsLogs_AgentEvent)(nil),
+17 -43
View File
@@ -3,62 +3,36 @@
package events
import (
"bytes"
"encoding/json"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"google.golang.org/protobuf/proto"
"github.com/ultravioletrs/cocos/agent/cvms"
)
type mockConn struct {
writeErr error
buf bytes.Buffer
}
func (m *mockConn) Write(p []byte) (n int, err error) {
if m.writeErr != nil {
return 0, m.writeErr
}
return m.buf.Write(p)
}
func TestSendEventSuccess(t *testing.T) {
mockConnection := &mockConn{}
svc, err := New("test_service", "12345", mockConnection)
queue := make(chan *cvms.ClientStreamMessage, 1)
svc, err := New("test_service", queue)
assert.NoError(t, err)
details := json.RawMessage(`{"key": "value"}`)
err = svc.SendEvent("test_event", "success", details)
assert.NoError(t, err)
go func() {
msg := <-queue
assert.NotNil(t, msg)
assert.NotNil(t, msg.GetAgentEvent())
assert.Equal(t, "test_event", msg.GetAgentEvent().EventType)
assert.Equal(t, "testid", msg.GetAgentEvent().ComputationId)
assert.Equal(t, "test_service", msg.GetAgentEvent().Originator)
assert.Equal(t, "success", msg.GetAgentEvent().Status)
var writtenMessage EventsLogs
err = proto.Unmarshal(mockConnection.buf.Bytes(), &writtenMessage)
assert.NoError(t, err)
now := time.Now()
eventTimestamp := msg.GetAgentEvent().GetTimestamp().AsTime()
assert.WithinDuration(t, now, eventTimestamp, 1*time.Second)
}()
assert.Equal(t, "test_event", writtenMessage.GetAgentEvent().EventType)
assert.Equal(t, "12345", writtenMessage.GetAgentEvent().ComputationId)
assert.Equal(t, "test_service", writtenMessage.GetAgentEvent().Originator)
assert.Equal(t, "success", writtenMessage.GetAgentEvent().Status)
svc.SendEvent("testid", "test_event", "success", details)
now := time.Now()
eventTimestamp := writtenMessage.GetAgentEvent().GetTimestamp().AsTime()
assert.WithinDuration(t, now, eventTimestamp, 1*time.Second)
}
func TestSendEventFailure(t *testing.T) {
mockConnection := &mockConn{writeErr: errors.New("write error")}
svc, err := New("test_service", "12345", mockConnection)
assert.NoError(t, err)
details := json.RawMessage(`{"key": "value"}`)
err = svc.SendEvent("test_event", "failure", details)
assert.Error(t, err)
assert.Equal(t, "write error", err.Error())
time.Sleep(1 * time.Second)
}
+11 -23
View File
@@ -24,22 +24,9 @@ func (_m *Service) EXPECT() *Service_Expecter {
return &Service_Expecter{mock: &_m.Mock}
}
// SendEvent provides a mock function with given fields: event, status, details
func (_m *Service) SendEvent(event string, status string, details json.RawMessage) error {
ret := _m.Called(event, status, details)
if len(ret) == 0 {
panic("no return value specified for SendEvent")
}
var r0 error
if rf, ok := ret.Get(0).(func(string, string, json.RawMessage) error); ok {
r0 = rf(event, status, details)
} else {
r0 = ret.Error(0)
}
return r0
// SendEvent provides a mock function with given fields: cmpID, event, status, details
func (_m *Service) SendEvent(cmpID string, event string, status string, details json.RawMessage) {
_m.Called(cmpID, event, status, details)
}
// Service_SendEvent_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendEvent'
@@ -48,26 +35,27 @@ type Service_SendEvent_Call struct {
}
// SendEvent is a helper method to define mock.On call
// - cmpID string
// - event string
// - status string
// - details json.RawMessage
func (_e *Service_Expecter) SendEvent(event interface{}, status interface{}, details interface{}) *Service_SendEvent_Call {
return &Service_SendEvent_Call{Call: _e.mock.On("SendEvent", event, status, details)}
func (_e *Service_Expecter) SendEvent(cmpID interface{}, event interface{}, status interface{}, details interface{}) *Service_SendEvent_Call {
return &Service_SendEvent_Call{Call: _e.mock.On("SendEvent", cmpID, event, status, details)}
}
func (_c *Service_SendEvent_Call) Run(run func(event string, status string, details json.RawMessage)) *Service_SendEvent_Call {
func (_c *Service_SendEvent_Call) Run(run func(cmpID string, event string, status string, details json.RawMessage)) *Service_SendEvent_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(string), args[1].(string), args[2].(json.RawMessage))
run(args[0].(string), args[1].(string), args[2].(string), args[3].(json.RawMessage))
})
return _c
}
func (_c *Service_SendEvent_Call) Return(_a0 error) *Service_SendEvent_Call {
_c.Call.Return(_a0)
func (_c *Service_SendEvent_Call) Return() *Service_SendEvent_Call {
_c.Call.Return()
return _c
}
func (_c *Service_SendEvent_Call) RunAndReturn(run func(string, string, json.RawMessage) error) *Service_SendEvent_Call {
func (_c *Service_SendEvent_Call) RunAndReturn(run func(string, string, string, json.RawMessage)) *Service_SendEvent_Call {
_c.Call.Return(run)
return _c
}
+93
View File
@@ -179,6 +179,53 @@ func (_c *Service_Data_Call) RunAndReturn(run func(context.Context, agent.Datase
return _c
}
// InitComputation provides a mock function with given fields: ctx, cmp
func (_m *Service) InitComputation(ctx context.Context, cmp agent.Computation) error {
ret := _m.Called(ctx, cmp)
if len(ret) == 0 {
panic("no return value specified for InitComputation")
}
var r0 error
if rf, ok := ret.Get(0).(func(context.Context, agent.Computation) error); ok {
r0 = rf(ctx, cmp)
} else {
r0 = ret.Error(0)
}
return r0
}
// Service_InitComputation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'InitComputation'
type Service_InitComputation_Call struct {
*mock.Call
}
// InitComputation is a helper method to define mock.On call
// - ctx context.Context
// - cmp agent.Computation
func (_e *Service_Expecter) InitComputation(ctx interface{}, cmp interface{}) *Service_InitComputation_Call {
return &Service_InitComputation_Call{Call: _e.mock.On("InitComputation", ctx, cmp)}
}
func (_c *Service_InitComputation_Call) Run(run func(ctx context.Context, cmp agent.Computation)) *Service_InitComputation_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context), args[1].(agent.Computation))
})
return _c
}
func (_c *Service_InitComputation_Call) Return(_a0 error) *Service_InitComputation_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *Service_InitComputation_Call) RunAndReturn(run func(context.Context, agent.Computation) error) *Service_InitComputation_Call {
_c.Call.Return(run)
return _c
}
// Result provides a mock function with given fields: ctx
func (_m *Service) Result(ctx context.Context) ([]byte, error) {
ret := _m.Called(ctx)
@@ -237,6 +284,52 @@ func (_c *Service_Result_Call) RunAndReturn(run func(context.Context) ([]byte, e
return _c
}
// StopComputation provides a mock function with given fields: ctx
func (_m *Service) StopComputation(ctx context.Context) error {
ret := _m.Called(ctx)
if len(ret) == 0 {
panic("no return value specified for StopComputation")
}
var r0 error
if rf, ok := ret.Get(0).(func(context.Context) error); ok {
r0 = rf(ctx)
} else {
r0 = ret.Error(0)
}
return r0
}
// Service_StopComputation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'StopComputation'
type Service_StopComputation_Call struct {
*mock.Call
}
// StopComputation is a helper method to define mock.On call
// - ctx context.Context
func (_e *Service_Expecter) StopComputation(ctx interface{}) *Service_StopComputation_Call {
return &Service_StopComputation_Call{Call: _e.mock.On("StopComputation", ctx)}
}
func (_c *Service_StopComputation_Call) Run(run func(ctx context.Context)) *Service_StopComputation_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context))
})
return _c
}
func (_c *Service_StopComputation_Call) Return(_a0 error) *Service_StopComputation_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *Service_StopComputation_Call) RunAndReturn(run func(context.Context) error) *Service_StopComputation_Call {
_c.Call.Return(run)
return _c
}
// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations.
// The first argument is typically a *testing.T value.
func NewService(t interface {
+68 -19
View File
@@ -104,6 +104,8 @@ var (
// Service specifies an API that must be fullfiled by the domain service
// implementation, and all of its decorators (e.g. logging & metrics).
type Service interface {
InitComputation(ctx context.Context, cmp Computation) error
StopComputation(ctx context.Context) error
Algo(ctx context.Context, algorithm Algorithm) error
Data(ctx context.Context, dataset Dataset) error
Result(ctx context.Context) ([]byte, error)
@@ -121,19 +123,21 @@ type agentService struct {
quoteProvider client.QuoteProvider // Provider for generating attestation quotes.
logger *slog.Logger // Logger for the agent service.
resultsConsumed bool // Indicates if the results have been consumed.
cancel context.CancelFunc // Cancels the computation context.
}
var _ Service = (*agentService)(nil)
// New instantiates the agent service implementation.
func New(ctx context.Context, logger *slog.Logger, eventSvc events.Service, cmp Computation, quoteProvider client.QuoteProvider) Service {
func New(ctx context.Context, logger *slog.Logger, eventSvc events.Service, quoteProvider client.QuoteProvider) Service {
sm := statemachine.NewStateMachine(Idle)
ctx, cancel := context.WithCancel(ctx)
svc := &agentService{
sm: sm,
eventSvc: eventSvc,
quoteProvider: quoteProvider,
logger: logger,
computation: cmp,
cancel: cancel,
}
transitions := []statemachine.Transition{
@@ -141,13 +145,6 @@ func New(ctx context.Context, logger *slog.Logger, eventSvc events.Service, cmp
{From: ReceivingManifest, Event: ManifestReceived, To: ReceivingAlgorithm},
}
if len(cmp.Datasets) == 0 {
transitions = append(transitions, statemachine.Transition{From: ReceivingAlgorithm, Event: AlgorithmReceived, To: Running})
} else {
transitions = append(transitions, statemachine.Transition{From: ReceivingAlgorithm, Event: AlgorithmReceived, To: ReceivingData})
transitions = append(transitions, statemachine.Transition{From: ReceivingData, Event: DataReceived, To: Running})
}
transitions = append(transitions, []statemachine.Transition{
{From: Running, Event: RunComplete, To: ConsumingResults},
{From: Running, Event: RunFailed, To: Failed},
@@ -158,8 +155,6 @@ func New(ctx context.Context, logger *slog.Logger, eventSvc events.Service, cmp
sm.AddTransition(t)
}
sm.SetAction(Idle, svc.publishEvent(IdleState.String()))
sm.SetAction(ReceivingManifest, svc.publishEvent(InProgress.String()))
sm.SetAction(ReceivingAlgorithm, svc.publishEvent(InProgress.String()))
sm.SetAction(ReceivingData, svc.publishEvent(InProgress.String()))
sm.SetAction(Running, svc.runComputation)
@@ -173,11 +168,67 @@ func New(ctx context.Context, logger *slog.Logger, eventSvc events.Service, cmp
}
}()
sm.SendEvent(Start)
defer sm.SendEvent(ManifestReceived)
return svc
}
func (as *agentService) InitComputation(ctx context.Context, cmp Computation) error {
defer as.sm.SendEvent(ManifestReceived)
if as.sm.GetState() != ReceivingManifest {
return ErrStateNotReady
}
as.mu.Lock()
defer as.mu.Unlock()
as.computation = cmp
transitions := []statemachine.Transition{}
if len(cmp.Datasets) == 0 {
transitions = append(transitions, statemachine.Transition{From: ReceivingAlgorithm, Event: AlgorithmReceived, To: Running})
} else {
transitions = append(transitions, statemachine.Transition{From: ReceivingAlgorithm, Event: AlgorithmReceived, To: ReceivingData})
transitions = append(transitions, statemachine.Transition{From: ReceivingData, Event: DataReceived, To: Running})
}
for _, t := range transitions {
as.sm.AddTransition(t)
}
return nil
}
func (as *agentService) StopComputation(ctx context.Context) error {
as.mu.Lock()
defer as.mu.Unlock()
as.cancel()
if err := as.algorithm.Stop(); err != nil {
return fmt.Errorf("error stopping computation: %v", err)
}
sm := statemachine.NewStateMachine(Idle)
if err := os.RemoveAll(algorithm.DatasetsDir); err != nil {
return fmt.Errorf("error removing datasets directory: %v", err)
}
if err := os.RemoveAll(algorithm.ResultsDir); err != nil {
return fmt.Errorf("error removing results directory: %v", err)
}
as.sm = sm
as.computation = Computation{}
as.algorithm = nil
as.result = nil
as.runError = nil
as.resultsConsumed = false
return nil
}
func (as *agentService) Algo(ctx context.Context, algo Algorithm) error {
if as.sm.GetState() != ReceivingAlgorithm {
return ErrStateNotReady
@@ -225,7 +276,7 @@ func (as *agentService) Algo(ctx context.Context, algo Algorithm) error {
switch algoType {
case string(algorithm.AlgoTypeBin):
as.algorithm = binary.NewAlgorithm(as.logger, as.eventSvc, f.Name(), args)
as.algorithm = binary.NewAlgorithm(as.logger, as.eventSvc, f.Name(), args, as.computation.ID)
case string(algorithm.AlgoTypePython):
var requirementsFile string
if len(algo.Requirements) > 0 {
@@ -243,11 +294,11 @@ func (as *agentService) Algo(ctx context.Context, algo Algorithm) error {
requirementsFile = fr.Name()
}
runtime := python.PythonRunTimeFromContext(ctx)
as.algorithm = python.NewAlgorithm(as.logger, as.eventSvc, runtime, requirementsFile, f.Name(), args)
as.algorithm = python.NewAlgorithm(as.logger, as.eventSvc, runtime, requirementsFile, f.Name(), args, as.computation.ID)
case string(algorithm.AlgoTypeWasm):
as.algorithm = wasm.NewAlgorithm(as.logger, as.eventSvc, f.Name(), args)
as.algorithm = wasm.NewAlgorithm(as.logger, as.eventSvc, args, f.Name(), as.computation.ID)
case string(algorithm.AlgoTypeDocker):
as.algorithm = docker.NewAlgorithm(as.logger, as.eventSvc, f.Name())
as.algorithm = docker.NewAlgorithm(as.logger, as.eventSvc, f.Name(), as.computation.ID)
}
if err := os.Mkdir(algorithm.DatasetsDir, 0o755); err != nil {
@@ -400,8 +451,6 @@ func (as *agentService) runComputation(state statemachine.State) {
func (as *agentService) publishEvent(status string) statemachine.Action {
return func(state statemachine.State) {
if err := as.eventSvc.SendEvent(state.String(), status, json.RawMessage{}); err != nil {
as.logger.Warn(err.Error())
}
as.eventSvc.SendEvent(as.computation.ID, state.String(), status, json.RawMessage{})
}
}
+23 -29
View File
@@ -35,11 +35,6 @@ var (
const datasetFile = "iris.csv"
func TestAlgo(t *testing.T) {
events := new(mocks.Service)
evCall := events.On("SendEvent", mock.Anything, mock.Anything, mock.Anything).Return(nil)
defer evCall.Unset()
qp, err := quoteprovider.GetQuoteProvider()
require.NoError(t, err)
@@ -120,9 +115,15 @@ func TestAlgo(t *testing.T) {
metadata.Pairs(algorithm.AlgoTypeKey, tc.algoType, python.PyRuntimeKey, python.PyRuntime),
)
events := new(mocks.Service)
events.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return()
ctx, cancel := context.WithCancel(ctx)
defer cancel()
svc := New(ctx, mglog.NewMock(), events, testComputation(t), qp)
svc := New(ctx, mglog.NewMock(), events, qp)
err := svc.InitComputation(ctx, testComputation(t))
require.NoError(t, err)
time.Sleep(300 * time.Millisecond)
@@ -138,11 +139,6 @@ func TestAlgo(t *testing.T) {
}
func TestData(t *testing.T) {
events := new(mocks.Service)
evCall := events.On("SendEvent", mock.Anything, mock.Anything, mock.Anything).Return(nil)
defer evCall.Unset()
qp, err := quoteprovider.GetQuoteProvider()
require.NoError(t, err)
@@ -209,6 +205,9 @@ func TestData(t *testing.T) {
python.PyRuntimeKey, python.PyRuntime),
)
events := new(mocks.Service)
events.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return()
if tc.err != ErrUndeclaredDataset {
ctx = IndexToContext(ctx, 0)
}
@@ -216,13 +215,16 @@ func TestData(t *testing.T) {
ctx, cancel := context.WithCancel(ctx)
defer cancel()
comp := testComputation(t)
svc := New(ctx, mglog.NewMock(), events, qp)
err := svc.InitComputation(ctx, testComputation(t))
require.NoError(t, err)
svc := New(ctx, mglog.NewMock(), events, comp, qp)
time.Sleep(300 * time.Millisecond)
if tc.err != ErrStateNotReady {
_ = svc.Algo(ctx, alg)
err = svc.Algo(ctx, alg)
require.NoError(t, err)
time.Sleep(300 * time.Millisecond)
}
err = svc.Data(ctx, tc.data)
@@ -238,11 +240,6 @@ func TestData(t *testing.T) {
}
func TestResult(t *testing.T) {
events := new(mocks.Service)
evCall := events.On("SendEvent", mock.Anything, mock.Anything, mock.Anything).Return(nil)
defer evCall.Unset()
qp, err := quoteprovider.GetQuoteProvider()
require.NoError(t, err)
@@ -285,6 +282,9 @@ func TestResult(t *testing.T) {
}
for _, tc := range cases {
events := new(mocks.Service)
events.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return()
t.Run(tc.name, func(t *testing.T) {
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(algorithm.AlgoTypeKey, "python", python.PyRuntimeKey, python.PyRuntime),
@@ -323,12 +323,8 @@ func TestResult(t *testing.T) {
}
func TestAttestation(t *testing.T) {
events := new(mocks.Service)
qp := new(mocks2.QuoteProvider)
evCall := events.On("SendEvent", mock.Anything, mock.Anything, mock.Anything).Return(nil)
defer evCall.Unset()
cases := []struct {
name string
reportData [ReportDataSize]byte
@@ -350,6 +346,9 @@ func TestAttestation(t *testing.T) {
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
events := new(mocks.Service)
events.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return()
ctx := metadata.NewIncomingContext(context.Background(),
metadata.Pairs(algorithm.AlgoTypeKey, "python", python.PyRuntimeKey, python.PyRuntime),
)
@@ -362,7 +361,7 @@ func TestAttestation(t *testing.T) {
}
defer getQuote.Unset()
svc := New(ctx, mglog.NewMock(), events, testComputation(t), qp)
svc := New(ctx, mglog.NewMock(), events, qp)
time.Sleep(300 * time.Millisecond)
_, err := svc.Attestation(ctx, tc.reportData)
assert.True(t, errors.Contains(err, tc.err), "expected %v, got %v", tc.err, err)
@@ -397,10 +396,5 @@ func testComputation(t *testing.T) Computation {
Datasets: []Dataset{{Hash: dataHash, UserKey: []byte("key"), Dataset: data, Filename: datasetFile}},
Algorithm: Algorithm{Hash: algoHash, UserKey: []byte("key"), Algorithm: algo},
ResultConsumers: []ResultConsumer{{UserKey: []byte("key")}},
AgentConfig: AgentConfig{
Port: "7002",
LogLevel: "debug",
AttestedTls: false,
},
}
}
+1 -1
View File
@@ -131,7 +131,7 @@ func TestManifestChecksum(t *testing.T) {
"name": "Example Computation",
"description": "This is an example computation"
}`,
expectedSum: "868825367c32c4b6d621d5d95e2890f233d8554df2348ab743aac2663a936f08",
expectedSum: "a99683e4d22ba54cefa51aa49fb2e97a92b828c088395992ddff16a6236f3299",
},
{
name: "Invalid JSON",
+130
View File
@@ -0,0 +1,130 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package cli
import (
"os"
"github.com/fatih/color"
"github.com/spf13/cobra"
"github.com/ultravioletrs/cocos/manager"
)
const (
serverURL = "server-url"
serverCA = "server-ca"
clientKey = "client-key"
clientCrt = "client-crt"
logLevel = "log-level"
)
var (
agentCVMServerUrl string
agentCVMServerCA string
agentCVMClientKey string
agentCVMClientCrt string
agentLogLevel string
)
func (c *CLI) NewCreateVMCmd() *cobra.Command {
cmd := &cobra.Command{
Use: "create-vm",
Short: "Create a new virtual machine",
Example: `create-vm`,
Args: cobra.ExactArgs(0),
Run: func(cmd *cobra.Command, args []string) {
if err := c.InitializeManagerClient(cmd); err != nil {
printError(cmd, "Failed to connect to manager: %v ❌ ", c.connectErr)
return
}
defer c.Close()
createReq, err := loadCerts()
if err != nil {
printError(cmd, "Error loading certs: %v ❌ ", err)
return
}
createReq.AgentCvmServerUrl = agentCVMServerUrl
createReq.AgentLogLevel = agentLogLevel
cmd.Println("🔗 Creating a new virtual machine")
res, err := c.managerClient.CreateVm(cmd.Context(), createReq)
if err != nil {
printError(cmd, "Error creating virtual machine: %v ❌ ", err)
return
}
cmd.Println(color.New(color.FgGreen).Sprintf("✅ Virtual machine created successfully with id %s and port %s", res.SvmId, res.ForwardedPort))
},
}
cmd.Flags().StringVar(&agentCVMServerUrl, serverURL, "", "CVM server URL")
cmd.Flags().StringVar(&agentCVMServerCA, serverCA, "", "CVM server CA")
cmd.Flags().StringVar(&agentCVMClientKey, clientKey, "", "CVM client key")
cmd.Flags().StringVar(&agentCVMClientCrt, clientCrt, "", "CVM client crt")
cmd.Flags().StringVar(&agentLogLevel, logLevel, "", "Agent Log level")
return cmd
}
func (c *CLI) NewRemoveVMCmd() *cobra.Command {
return &cobra.Command{
Use: "remove-vm",
Short: "Remove a virtual machine",
Example: `remove-vm <svm_id>`,
Args: cobra.ExactArgs(1),
Run: func(cmd *cobra.Command, args []string) {
if err := c.InitializeManagerClient(cmd); err == nil {
defer c.Close()
}
if c.connectErr != nil {
printError(cmd, "Failed to connect to manager: %v ❌ ", c.connectErr)
return
}
cmd.Println("🔗 Removing virtual machine")
_, err := c.managerClient.RemoveVm(cmd.Context(), &manager.RemoveReq{SvmId: args[0]})
if err != nil {
printError(cmd, "Error removing virtual machine: %v ❌ ", err)
return
}
cmd.Println(color.New(color.FgGreen).Sprintf("✅ Virtual machine removed successfully"))
},
}
}
func fileReader(path string) ([]byte, error) {
if path == "" {
return nil, nil
}
return os.ReadFile(path)
}
func loadCerts() (*manager.CreateReq, error) {
clientKey, err := fileReader(agentCVMClientKey)
if err != nil {
return nil, err
}
clientCrt, err := fileReader(agentCVMClientCrt)
if err != nil {
return nil, err
}
serverCA, err := fileReader(agentCVMServerCA)
if err != nil {
return nil, err
}
return &manager.CreateReq{
AgentCvmServerCaCert: serverCA,
AgentCvmClientKey: clientKey,
AgentCvmClientCert: clientCrt,
}, nil
}
+27 -8
View File
@@ -6,28 +6,33 @@ import (
"context"
"github.com/spf13/cobra"
"github.com/ultravioletrs/cocos/manager"
"github.com/ultravioletrs/cocos/pkg/clients/grpc"
"github.com/ultravioletrs/cocos/pkg/clients/grpc/agent"
managergrpc "github.com/ultravioletrs/cocos/pkg/clients/grpc/manager"
"github.com/ultravioletrs/cocos/pkg/sdk"
)
var Verbose bool
type CLI struct {
agentSDK sdk.SDK
config grpc.AgentClientConfig
client grpc.Client
connectErr error
agentSDK sdk.SDK
agentConfig grpc.AgentClientConfig
managerConfig grpc.ManagerClientConfig
client grpc.Client
managerClient manager.ManagerServiceClient
connectErr error
}
func New(config grpc.AgentClientConfig) *CLI {
func New(agentConfig grpc.AgentClientConfig, managerConfig grpc.ManagerClientConfig) *CLI {
return &CLI{
config: config,
agentConfig: agentConfig,
managerConfig: managerConfig,
}
}
func (c *CLI) InitializeSDK(cmd *cobra.Command) error {
agentGRPCClient, agentClient, err := agent.NewAgentClient(context.Background(), c.config)
func (c *CLI) InitializeAgentSDK(cmd *cobra.Command) error {
agentGRPCClient, agentClient, err := agent.NewAgentClient(context.Background(), c.agentConfig)
if err != nil {
c.connectErr = err
return err
@@ -39,6 +44,20 @@ func (c *CLI) InitializeSDK(cmd *cobra.Command) error {
return nil
}
func (c *CLI) InitializeManagerClient(cmd *cobra.Command) error {
managerGRPCClient, managerClient, err := managergrpc.NewManagerClient(c.managerConfig)
if err != nil {
c.connectErr = err
return err
}
cmd.Println("🔗 Connected to manager using ", managerGRPCClient.Secure())
c.client = managerGRPCClient
c.managerClient = managerClient
return nil
}
func (c *CLI) Close() {
c.client.Close()
}
+57 -182
View File
@@ -3,77 +3,68 @@
package main
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log"
"log/slog"
"os"
"os/signal"
"syscall"
"time"
mglog "github.com/absmach/magistrala/logger"
"github.com/absmach/magistrala/pkg/prometheus"
"github.com/cenkalti/backoff/v4"
"github.com/google/go-sev-guest/abi"
"github.com/caarlos0/env/v11"
"github.com/google/go-sev-guest/client"
"github.com/mdlayher/vsock"
"github.com/ultravioletrs/cocos/agent"
"github.com/ultravioletrs/cocos/agent/api"
agentgrpc "github.com/ultravioletrs/cocos/agent/api/grpc"
"github.com/ultravioletrs/cocos/agent/auth"
"github.com/ultravioletrs/cocos/agent/cvms"
cvmapi "github.com/ultravioletrs/cocos/agent/cvms/api/grpc"
"github.com/ultravioletrs/cocos/agent/cvms/server"
"github.com/ultravioletrs/cocos/agent/events"
agentlogger "github.com/ultravioletrs/cocos/internal/logger"
"github.com/ultravioletrs/cocos/internal/server"
grpcserver "github.com/ultravioletrs/cocos/internal/server/grpc"
ackvsock "github.com/ultravioletrs/cocos/internal/vsock"
managerevents "github.com/ultravioletrs/cocos/manager/events"
"github.com/ultravioletrs/cocos/manager/qemu"
"github.com/ultravioletrs/cocos/pkg/attestation/quoteprovider"
"golang.org/x/crypto/sha3"
pkggrpc "github.com/ultravioletrs/cocos/pkg/clients/grpc"
cvmgrpc "github.com/ultravioletrs/cocos/pkg/clients/grpc/cvm"
"golang.org/x/sync/errgroup"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
)
const (
svcName = "agent"
defSvcGRPCPort = "7002"
retryInterval = 5 * time.Second
svcName = "agent"
defSvcGRPCPort = "7002"
retryInterval = 5 * time.Second
envPrefixCVMGRPC = "AGENT_CVM_GRPC_"
)
type config struct {
LogLevel string `env:"AGENT_LOG_LEVEL" envDefault:"debug"`
}
func main() {
ctx, cancel := context.WithCancel(context.Background())
g, ctx := errgroup.WithContext(ctx)
cfg, err := readConfig()
if err != nil {
log.Fatalf("failed to read agent configuration from vsock %s", err.Error())
var cfg config
if err := env.Parse(&cfg); err != nil {
log.Fatalf("failed to load %s configuration : %s", svcName, err)
}
conn, err := dialVsock()
if err != nil {
log.Fatal(err)
}
defer conn.Close()
ackConn := ackvsock.NewAckWriter(conn)
var exitCode int
defer mglog.ExitWithError(&exitCode)
var level slog.Level
if err := level.UnmarshalText([]byte(cfg.AgentConfig.LogLevel)); err != nil {
if err := level.UnmarshalText([]byte(cfg.LogLevel)); err != nil {
log.Println(err)
exitCode = 1
return
}
handler := agentlogger.NewProtoHandler(ackConn, &slog.HandlerOptions{Level: level}, cfg.ID)
eventsLogsQueue := make(chan *cvms.ClientStreamMessage, 1000)
handler := agentlogger.NewProtoHandler(os.Stdout, &slog.HandlerOptions{Level: level}, eventsLogsQueue)
logger := slog.New(handler)
eventSvc, err := events.New(svcName, cfg.ID, ackConn)
eventSvc, err := events.New(svcName, eventsLogsQueue)
if err != nil {
logger.Error(fmt.Sprintf("failed to create events service %s", err.Error()))
exitCode = 1
@@ -87,64 +78,49 @@ func main() {
return
}
if err := verifyManifest(cfg, qp); err != nil {
cvmGrpcConfig := pkggrpc.CVMClientConfig{}
if err := env.ParseWithOptions(&cvmGrpcConfig, env.Options{Prefix: envPrefixCVMGRPC}); err != nil {
logger.Error(fmt.Sprintf("failed to load %s gRPC client configuration : %s", svcName, err))
exitCode = 1
return
}
cvmGRPCClient, cvmClient, err := cvmgrpc.NewCVMClient(cvmGrpcConfig)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
}
defer cvmGRPCClient.Close()
pc, err := cvmClient.Process(ctx)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
}
setDefaultValues(&cfg)
svc := newService(ctx, logger, eventSvc, qp)
svc := newService(ctx, logger, eventSvc, cfg, qp)
agentGrpcServerConfig := server.AgentConfig{
ServerConfig: server.ServerConfig{
BaseConfig: server.BaseConfig{
Host: cfg.AgentConfig.Host,
Port: cfg.AgentConfig.Port,
CertFile: cfg.AgentConfig.CertFile,
KeyFile: cfg.AgentConfig.KeyFile,
ServerCAFile: cfg.AgentConfig.ServerCAFile,
ClientCAFile: cfg.AgentConfig.ClientCAFile,
},
},
AttestedTLS: cfg.AgentConfig.AttestedTls,
}
registerAgentServiceServer := func(srv *grpc.Server) {
reflection.Register(srv)
agent.RegisterAgentServiceServer(srv, agentgrpc.NewServer(svc))
}
authSvc, err := auth.New(cfg)
if err != nil {
logger.Error(fmt.Sprintf("failed to create auth service %s", err.Error()))
exitCode = 1
return
}
gs := grpcserver.New(ctx, cancel, svcName, agentGrpcServerConfig, registerAgentServiceServer, logger, qp, authSvc)
mc := cvmapi.NewClient(pc, svc, eventsLogsQueue, logger, server.NewServer(logger, svc))
g.Go(func() error {
for {
if _, err := io.Copy(io.Discard, conn); err != nil {
log.Printf("vsock connection lost: %v, reconnecting...", err)
conn.Close()
conn, err = dialVsock()
if err != nil {
log.Fatal("failed to reconnect: ", err)
}
}
time.Sleep(retryInterval)
ch := make(chan os.Signal, 1)
signal.Notify(ch, syscall.SIGINT, syscall.SIGTERM)
defer signal.Stop(ch)
select {
case <-ch:
logger.Info("Received signal, shutting down...")
cancel()
return nil
case <-ctx.Done():
return ctx.Err()
}
})
g.Go(func() error {
return gs.Start()
})
g.Go(func() error {
return server.StopHandler(ctx, cancel, logger, svcName, gs)
return mc.Process(ctx, cancel)
})
if err := g.Wait(); err != nil {
@@ -152,8 +128,8 @@ func main() {
}
}
func newService(ctx context.Context, logger *slog.Logger, eventSvc events.Service, cmp agent.Computation, qp client.QuoteProvider) agent.Service {
svc := agent.New(ctx, logger, eventSvc, cmp, qp)
func newService(ctx context.Context, logger *slog.Logger, eventSvc events.Service, qp client.QuoteProvider) agent.Service {
svc := agent.New(ctx, logger, eventSvc, qp)
svc = api.LoggingMiddleware(svc, logger)
counter, latency := prometheus.MakeMetrics(svcName, "api")
@@ -161,104 +137,3 @@ func newService(ctx context.Context, logger *slog.Logger, eventSvc events.Servic
return svc
}
func readConfig() (agent.Computation, error) {
l, err := vsock.Listen(qemu.VsockConfigPort, nil)
if err != nil {
return agent.Computation{}, err
}
defer l.Close()
conn, err := l.Accept()
if err != nil {
return agent.Computation{}, err
}
defer conn.Close()
var buffer []byte
for {
chunk := make([]byte, 1024)
n, err := conn.Read(chunk)
if err != nil {
if err == io.EOF {
break
}
return agent.Computation{}, err
}
buffer = append(buffer, chunk[:n]...)
}
ac := agent.Computation{
AgentConfig: agent.AgentConfig{},
}
if err := json.Unmarshal(buffer, &ac); err != nil {
return agent.Computation{}, err
}
return ac, nil
}
func setDefaultValues(cfg *agent.Computation) {
if cfg.AgentConfig.LogLevel == "" {
cfg.AgentConfig.LogLevel = "info"
}
if cfg.AgentConfig.Port == "" {
cfg.AgentConfig.Port = defSvcGRPCPort
}
}
func isTEE() bool {
_, err := os.Stat("/dev/sev-guest")
return !os.IsNotExist(err)
}
func dialVsock() (*vsock.Conn, error) {
var conn *vsock.Conn
var err error
err = backoff.Retry(func() error {
conn, err = vsock.Dial(vsock.Host, managerevents.ManagerVsockPort, nil)
if err == nil {
log.Println("vsock connection established")
return nil
}
log.Printf("vsock connection failed, retrying in %s... Error: %v", retryInterval, err)
return err
}, backoff.NewExponentialBackOff())
if err != nil {
return nil, err
}
return conn, nil
}
func verifyManifest(cfg agent.Computation, qp client.QuoteProvider) error {
if !isTEE() {
return nil
}
ar, err := qp.GetRawQuote(sha3.Sum512([]byte(cfg.ID)))
if err != nil {
return err
}
arProto, err := abi.ReportCertsToProto(ar[:abi.ReportSize])
if err != nil {
return err
}
cfgBytes, err := json.Marshal(cfg)
if err != nil {
return err
}
mcHash := sha3.Sum256(cfgBytes)
if arProto.Report.HostData == nil {
return fmt.Errorf("manifest verification failed: HostData is nil")
}
if !bytes.Equal(arProto.Report.HostData, mcHash[:]) {
return fmt.Errorf("manifest verification failed")
}
return nil
}
-94
View File
@@ -1,94 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package main
import (
"context"
"log/slog"
"os"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/ultravioletrs/cocos/agent"
"github.com/ultravioletrs/cocos/agent/events/mocks"
qpmocks "github.com/ultravioletrs/cocos/pkg/attestation/quoteprovider/mocks"
)
func TestSetDefaultValues(t *testing.T) {
tests := []struct {
name string
input agent.Computation
expected agent.Computation
}{
{
name: "Empty config",
input: agent.Computation{
AgentConfig: agent.AgentConfig{},
},
expected: agent.Computation{
AgentConfig: agent.AgentConfig{
LogLevel: "info",
Port: "7002",
},
},
},
{
name: "Partial config",
input: agent.Computation{
AgentConfig: agent.AgentConfig{
LogLevel: "debug",
},
},
expected: agent.Computation{
AgentConfig: agent.AgentConfig{
LogLevel: "debug",
Port: "7002",
},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
setDefaultValues(&tt.input)
assert.Equal(t, tt.expected, tt.input)
})
}
}
func TestNewService(t *testing.T) {
ctx := context.Background()
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
eventSvc := new(mocks.Service)
eventSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything).Return(nil)
cmp := agent.Computation{
ID: "test-computation",
AgentConfig: agent.AgentConfig{
LogLevel: "info",
Port: "7002",
},
}
qp := new(qpmocks.QuoteProvider)
svc := newService(ctx, logger, eventSvc, cmp, qp)
assert.NotNil(t, svc)
}
func TestVerifyManifest(t *testing.T) {
cfg := agent.Computation{
ID: "test-computation",
AgentConfig: agent.AgentConfig{
LogLevel: "info",
Port: "7002",
},
}
mockQP := new(qpmocks.QuoteProvider)
mockQP.On("GetRawQuote", mock.Anything).Return([]byte{}, nil)
err := verifyManifest(cfg, mockQP)
assert.NoError(t, err)
}
+17 -7
View File
@@ -19,11 +19,12 @@ import (
)
const (
svcName = "cli"
envPrefixAgentGRPC = "AGENT_GRPC_"
completion = "completion"
filePermision = 0o755
cocosDirectory = ".cocos"
svcName = "cli"
envPrefixAgentGRPC = "AGENT_GRPC_"
envPrefixManagerGRPC = "MANAGER_GRPC_"
completion = "completion"
filePermision = 0o755
cocosDirectory = ".cocos"
)
type config struct {
@@ -98,9 +99,16 @@ func main() {
return
}
cliSVC := cli.New(agentGRPCConfig)
managerGRPCConfig := grpc.ManagerClientConfig{}
if err := env.ParseWithOptions(&managerGRPCConfig, env.Options{Prefix: envPrefixManagerGRPC}); err != nil {
message := color.New(color.FgRed).Sprintf("failed to load %s gRPC client configuration : %s", svcName, err)
rootCmd.Println(message)
return
}
if err := cliSVC.InitializeSDK(rootCmd); err == nil {
cliSVC := cli.New(agentGRPCConfig, managerGRPCConfig)
if err := cliSVC.InitializeAgentSDK(rootCmd); err == nil {
defer cliSVC.Close()
}
@@ -119,6 +127,8 @@ func main() {
rootCmd.AddCommand(attestationPolicyCmd)
rootCmd.AddCommand(keysCmd)
rootCmd.AddCommand(cliSVC.NewCABundleCmd(directoryCachePath))
rootCmd.AddCommand(cliSVC.NewCreateVMCmd())
rootCmd.AddCommand(cliSVC.NewRemoveVMCmd())
// Attestation commands
attestationCmd.AddCommand(cliSVC.NewGetAttestationCmd())
+15 -47
View File
@@ -10,25 +10,24 @@ import (
"log/slog"
"net/url"
"os"
"os/signal"
"strings"
"syscall"
mglog "github.com/absmach/magistrala/logger"
"github.com/absmach/magistrala/pkg/jaeger"
"github.com/absmach/magistrala/pkg/prometheus"
"github.com/absmach/magistrala/pkg/uuid"
"github.com/caarlos0/env/v11"
"github.com/ultravioletrs/cocos/internal/server"
grpcserver "github.com/ultravioletrs/cocos/internal/server/grpc"
"github.com/ultravioletrs/cocos/manager"
"github.com/ultravioletrs/cocos/manager/api"
managerapi "github.com/ultravioletrs/cocos/manager/api/grpc"
"github.com/ultravioletrs/cocos/manager/events"
managergrpc "github.com/ultravioletrs/cocos/manager/api/grpc"
"github.com/ultravioletrs/cocos/manager/qemu"
"github.com/ultravioletrs/cocos/manager/tracing"
pkggrpc "github.com/ultravioletrs/cocos/pkg/clients/grpc"
managergrpc "github.com/ultravioletrs/cocos/pkg/clients/grpc/manager"
"go.opentelemetry.io/otel/trace"
"golang.org/x/sync/errgroup"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection"
)
const (
@@ -92,64 +91,33 @@ func main() {
args := qemuCfg.ConstructQemuArgs()
logger.Info(strings.Join(args, " "))
managerGRPCConfig := pkggrpc.ManagerClientConfig{}
managerGRPCConfig := server.ServerConfig{}
if err := env.ParseWithOptions(&managerGRPCConfig, env.Options{Prefix: envPrefixGRPC}); err != nil {
logger.Error(fmt.Sprintf("failed to load %s gRPC client configuration : %s", svcName, err))
exitCode = 1
return
}
managerGRPCClient, managerClient, err := managergrpc.NewManagerClient(managerGRPCConfig)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
}
defer managerGRPCClient.Close()
pc, err := managerClient.Process(ctx)
svc, err := newService(logger, tracer, qemuCfg, cfg.AttestationPolicyBinary, cfg.EosVersion)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
}
eventsChan := make(chan *manager.ClientStreamMessage, clientBufferSize)
svc, err := newService(logger, tracer, qemuCfg, eventsChan, cfg.AttestationPolicyBinary, cfg.EosVersion)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
registerManagerServiceServer := func(srv *grpc.Server) {
reflection.Register(srv)
manager.RegisterManagerServiceServer(srv, managergrpc.NewServer(svc))
}
eventsSvc, err := events.New(logger, svc.ReportBrokenConnection, eventsChan)
if err != nil {
logger.Error(err.Error())
exitCode = 1
return
}
go eventsSvc.Listen(ctx)
mc := managerapi.NewClient(pc, svc, eventsChan, logger)
gs := grpcserver.New(ctx, cancel, svcName, managerGRPCConfig, registerManagerServiceServer, logger, nil, nil)
g.Go(func() error {
ch := make(chan os.Signal, 1)
signal.Notify(ch, syscall.SIGINT, syscall.SIGTERM)
defer signal.Stop(ch)
select {
case <-ch:
logger.Info("Received signal, shutting down...")
cancel()
return nil
case <-ctx.Done():
return ctx.Err()
}
return gs.Start()
})
g.Go(func() error {
return mc.Process(ctx, cancel)
return server.StopHandler(ctx, cancel, logger, svcName, gs)
})
if err := g.Wait(); err != nil {
@@ -157,8 +125,8 @@ func main() {
}
}
func newService(logger *slog.Logger, tracer trace.Tracer, qemuCfg qemu.Config, eventsChan chan *manager.ClientStreamMessage, attestationPolicyPath string, eosVersion string) (manager.Service, error) {
svc, err := manager.New(qemuCfg, attestationPolicyPath, logger, eventsChan, qemu.NewVM, eosVersion)
func newService(logger *slog.Logger, tracer trace.Tracer, qemuCfg qemu.Config, attestationPolicyPath string, eosVersion string) (manager.Service, error) {
svc, err := manager.New(qemuCfg, attestationPolicyPath, logger, qemu.NewVM, eosVersion)
if err != nil {
return nil, err
}
+6 -5
View File
@@ -10,7 +10,8 @@ MANAGER_ATTESTATION_POLICY_BINARY=../../build
MANAGER_GRPC_CLIENT_CERT=
MANAGER_GRPC_CLIENT_KEY=
MANAGER_GRPC_SERVER_CA_CERTS=
MANAGER_GRPC_URL=localhost:7001
MANAGER_GRPC_PORT=6101
MANAGER_GRPC_HOST=0.0.0.0
MANAGER_GRPC_TIMEOUT=60s
MANAGER_EOS_VERSION=""
@@ -21,13 +22,13 @@ MANAGER_QEMU_MAX_MEMORY=30G
MANAGER_QEMU_OVMF_CODE_IF=pflash
MANAGER_QEMU_OVMF_CODE_FORMAT=raw
MANAGER_QEMU_OVMF_CODE_UNIT=0
MANAGER_QEMU_OVMF_CODE_FILE=/usr/share/OVMF/x64/OVMF_CODE.fd
MANAGER_QEMU_OVMF_CODE_FILE=/usr/share/edk2/x64/OVMF_CODE.fd
MANAGER_QEMU_OVMF_VERSION=edk2-stable202408
MANAGER_QEMU_OVMF_CODE_READONLY=on
MANAGER_QEMU_OVMF_VARS_IF=pflash
MANAGER_QEMU_OVMF_VARS_FORMAT=raw
MANAGER_QEMU_OVMF_VARS_UNIT=1
MANAGER_QEMU_OVMF_VARS_FILE=/usr/share/OVMF/x64/OVMF_VARS.fd
MANAGER_QEMU_OVMF_VARS_FILE=/usr/share/edk2/x64/OVMF_VARS.fd
MANAGER_QEMU_NETDEV_ID=vmnic
MANAGER_QEMU_HOST_FWD_AGENT=7020
MANAGER_QEMU_GUEST_FWD_AGENT=7002
@@ -35,8 +36,8 @@ MANAGER_QEMU_VIRTIO_NET_PCI_DISABLE_LEGACY=on
MANAGER_QEMU_VIRTIO_NET_PCI_IOMMU_PLATFORM=true
MANAGER_QEMU_VIRTIO_NET_PCI_ADDR=0x2
MANAGER_QEMU_VIRTIO_NET_PCI_ROMFILE=
MANAGER_QEMU_DISK_IMG_KERNEL_FILE=/home/sammyk/Documents/cocos-ai/cmd/manager/img/bzImage
MANAGER_QEMU_DISK_IMG_ROOTFS_FILE=/home/sammyk/Documents/cocos-ai/cmd/manager/img/rootfs.cpio.gz
MANAGER_QEMU_DISK_IMG_KERNEL_FILE=/etc/cocos/bzImage
MANAGER_QEMU_DISK_IMG_ROOTFS_FILE=/etc/cocos/rootfs.cpio.gz
MANAGER_QEMU_SEV_ID=sev0
MANAGER_QEMU_SEV_CBITPOS=51
MANAGER_QEMU_SEV_REDUCED_PHYS_BITS=1
+36 -35
View File
@@ -3,47 +3,50 @@ module github.com/ultravioletrs/cocos
go 1.23.0
require (
github.com/absmach/magistrala v0.14.1-0.20240709113739-04c359462746
github.com/caarlos0/env/v11 v11.2.2
github.com/cenkalti/backoff/v4 v4.3.0
github.com/absmach/magistrala v0.15.1
github.com/caarlos0/env/v11 v11.3.1
github.com/fatih/color v1.18.0
github.com/go-kit/kit v0.13.0
github.com/gofrs/uuid v4.4.0+incompatible
github.com/google/go-sev-guest v0.11.1
github.com/mdlayher/vsock v1.2.1
github.com/spf13/cobra v1.8.1
github.com/spf13/pflag v1.0.5
github.com/stretchr/testify v1.9.0
github.com/spf13/pflag v1.0.6
github.com/stretchr/testify v1.10.0
github.com/virtee/sev-snp-measure-go v0.0.0-20240530153610-e6e8dc9b6877
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0
go.opentelemetry.io/otel/trace v1.32.0
golang.org/x/crypto v0.29.0
golang.org/x/sync v0.9.0
google.golang.org/grpc v1.68.0
google.golang.org/protobuf v1.35.2
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0
go.opentelemetry.io/otel/trace v1.34.0
golang.org/x/crypto v0.32.0
golang.org/x/sync v0.10.0
google.golang.org/grpc v1.69.4
google.golang.org/protobuf v1.36.3
)
require (
github.com/Microsoft/go-winio v0.6.1 // indirect
github.com/Microsoft/go-winio v0.6.2 // indirect
github.com/cenkalti/backoff/v4 v4.3.0 // indirect
github.com/containerd/log v0.1.0 // indirect
github.com/distribution/reference v0.6.0 // indirect
github.com/docker/go-connections v0.5.0 // indirect
github.com/docker/go-units v0.5.0 // indirect
github.com/felixge/httpsnoop v1.0.4 // indirect
github.com/gofrs/uuid/v5 v5.3.0 // indirect
github.com/gogo/protobuf v1.3.2 // indirect
github.com/golang/protobuf v1.5.4 // indirect
github.com/mattn/go-colorable v0.1.13 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/moby/docker-image-spec v1.3.1 // indirect
github.com/morikuni/aec v1.0.0 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/opencontainers/go-digest v1.0.0 // indirect
github.com/opencontainers/image-spec v1.1.0 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.53.0 // indirect
go.opentelemetry.io/otel v1.32.0 // indirect
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.28.0 // indirect
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.28.0 // indirect
go.opentelemetry.io/otel/sdk v1.28.0 // indirect
golang.org/x/mod v0.19.0 // indirect
golang.org/x/tools v0.23.0 // indirect
github.com/pborman/uuid v1.2.1 // indirect
go.opentelemetry.io/auto/sdk v1.1.0 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0 // indirect
go.opentelemetry.io/otel v1.34.0 // indirect
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.32.0 // indirect
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.32.0 // indirect
go.opentelemetry.io/otel/sdk v1.32.0 // indirect
gotest.tools/v3 v3.5.1 // indirect
)
@@ -51,36 +54,34 @@ require (
github.com/beorn7/perks v1.0.1 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
github.com/docker/docker v27.3.1+incompatible
github.com/docker/docker v27.5.1+incompatible
github.com/go-kit/log v0.2.1 // indirect
github.com/go-logfmt/logfmt v0.6.0 // indirect
github.com/go-logr/logr v1.4.2 // indirect
github.com/go-logr/stdr v1.2.2 // indirect
github.com/golang/protobuf v1.5.4 // indirect
github.com/google/go-configfs-tsm v0.2.2 // indirect
github.com/google/logger v1.1.1
github.com/google/uuid v1.6.0 // indirect
github.com/grpc-ecosystem/grpc-gateway/v2 v2.20.0 // indirect
github.com/google/uuid v1.6.0
github.com/grpc-ecosystem/grpc-gateway/v2 v2.23.0 // indirect
github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/mdlayher/socket v0.4.1 // indirect
github.com/pborman/uuid v1.2.1 // indirect
github.com/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect
github.com/prometheus/client_golang v1.19.1 // indirect
github.com/prometheus/client_golang v1.20.5 // indirect
github.com/prometheus/client_model v0.6.1 // indirect
github.com/prometheus/common v0.52.2 // indirect
github.com/prometheus/procfs v0.13.0 // indirect
github.com/prometheus/common v0.59.1 // indirect
github.com/prometheus/procfs v0.15.1 // indirect
github.com/stretchr/objx v0.5.2 // indirect
go.opentelemetry.io/otel/metric v1.32.0 // indirect
go.opentelemetry.io/otel/metric v1.34.0 // indirect
go.opentelemetry.io/proto/otlp v1.3.1 // indirect
go.uber.org/multierr v1.11.0 // indirect
golang.org/x/net v0.30.0 // indirect
golang.org/x/sys v0.27.0 // indirect
golang.org/x/term v0.26.0
golang.org/x/text v0.20.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20240903143218-8af14fe29dc1 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28 // indirect
golang.org/x/net v0.34.0 // indirect
golang.org/x/sys v0.29.0 // indirect
golang.org/x/term v0.28.0
golang.org/x/text v0.21.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20241104194629-dd2ea8efbc28 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250115164207-1a7da9e5054f // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
)
replace github.com/virtee/sev-snp-measure-go => github.com/sammyoina/sev-snp-measure-go v0.0.0-20241107163739-38915ab517c7
replace github.com/virtee/sev-snp-measure-go => github.com/sammyoina/sev-snp-measure-go v0.0.0-20241202151803-ef189f0ff825
+72 -65
View File
@@ -1,15 +1,15 @@
github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161 h1:L/gRVlceqvL25UVaW/CKtUDjefjrs0SPonmDGUVOYP0=
github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
github.com/Microsoft/go-winio v0.6.1 h1:9/kr64B9VUZrLm5YYwbGtUJnMgqWVOdUAXu6Migciow=
github.com/Microsoft/go-winio v0.6.1/go.mod h1:LRdKpFKfdobln8UmuiYcKPot9D2v6svN5+sAH+4kjUM=
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
github.com/VividCortex/gohistogram v1.0.0 h1:6+hBz+qvs0JOrrNhhmR7lFxo5sINxBCGXrdtl/UvroE=
github.com/VividCortex/gohistogram v1.0.0/go.mod h1:Pf5mBqqDxYaXu3hDrrU+w6nw50o/4+TcAqDqk/vUH7g=
github.com/absmach/magistrala v0.14.1-0.20240709113739-04c359462746 h1:Tj567KeGVygjTsSCxn4++skKiz9GkPugM1KMdIFxvfw=
github.com/absmach/magistrala v0.14.1-0.20240709113739-04c359462746/go.mod h1:CIx3OsPFc4doJZmBWSA6LNWefcznKv9c3cLOxNxL4q4=
github.com/absmach/magistrala v0.15.1 h1:3Bk2hlyWcV591LxPYwlvRcyCXTfuZ1g/EkNmU+o3NNQ=
github.com/absmach/magistrala v0.15.1/go.mod h1:9pto6xuBt/IuCtZRdEha0iDQKNQ5tyNOjLXJgUiikYk=
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
github.com/caarlos0/env/v11 v11.2.2 h1:95fApNrUyueipoZN/EhA8mMxiNxrBwDa+oAZrMWl3Kg=
github.com/caarlos0/env/v11 v11.2.2/go.mod h1:JBfcdeQiBoI3Zh1QRAWfe+tpiNTmDtcCj/hHHHMx0vc=
github.com/caarlos0/env/v11 v11.3.1 h1:cArPWC15hWmEt+gWk7YBi7lEXTXCvpaSdCiZE2X5mCA=
github.com/caarlos0/env/v11 v11.3.1/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U=
github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8=
github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
@@ -21,8 +21,8 @@ github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
github.com/docker/docker v27.3.1+incompatible h1:KttF0XoteNTicmUtBO0L2tP+J7FGRFTjaEF4k6WdhfI=
github.com/docker/docker v27.3.1+incompatible/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk=
github.com/docker/docker v27.5.1+incompatible h1:4PYU5dnBYqRQi0294d1FBECqT9ECWeQAIfE8q4YnPY8=
github.com/docker/docker v27.5.1+incompatible/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk=
github.com/docker/go-connections v0.5.0 h1:USnMq7hx7gwdVZq1L49hLXaFtUdTADjXGp+uj1Br63c=
github.com/docker/go-connections v0.5.0/go.mod h1:ov60Kzw0kKElRwhNs9UlUHAE/F9Fe6GLaXnqyDdmEXc=
github.com/docker/go-units v0.5.0 h1:69rxXcBk27SvSaaxTtLh/8llcHD8vYHT7WSdRZ/jvr4=
@@ -44,6 +44,8 @@ github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/gofrs/uuid v4.4.0+incompatible h1:3qXRTX8/NbyulANqlc0lchS1gqAVxRgsuW1YrTJupqA=
github.com/gofrs/uuid v4.4.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRxnplIgP/c0N/04lM=
github.com/gofrs/uuid/v5 v5.3.0 h1:m0mUMr+oVYUdxpMLgSYCZiXe7PuVPnI94+OMeVBNedk=
github.com/gofrs/uuid/v5 v5.3.0/go.mod h1:CDOjlDMVAtN56jqyRUZh58JT31Tiw7/oQyEXZV+9bD8=
github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q=
github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
@@ -59,12 +61,14 @@ github.com/google/logger v1.1.1/go.mod h1:BkeJZ+1FhQ+/d087r4dzojEg1u2ZX+ZqG1jTUr
github.com/google/uuid v1.0.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.20.0 h1:bkypFPDjIYGfCYD5mRBvpqxfYX1YCS1PXdKYWi8FsN0=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.20.0/go.mod h1:P+Lt/0by1T8bfcF3z737NnSbmxQAppXMRziHUxPOC8k=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.23.0 h1:ad0vkEBuk23VJzZR9nkLVG0YAoN9coASF1GusYX6AlU=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.23.0/go.mod h1:igFoXX2ELCW06bol23DWPB5BEWfZISOzSP5K2sbLea0=
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
github.com/klauspost/compress v1.17.9 h1:6KIumPrER1LHsvBVuDa0r5xaG0Es51mhhB9BQB2qeMA=
github.com/klauspost/compress v1.17.9/go.mod h1:Di0epgTjJY877eYKx5yC51cX2A2Vl2ibi7bDH9ttBbw=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
@@ -84,6 +88,8 @@ github.com/moby/term v0.5.0 h1:xt8Q1nalod/v7BqbG21f8mQPqH+xAaC9C3N3wfWbVP0=
github.com/moby/term v0.5.0/go.mod h1:8FzsFHVUBGZdbDsJw/ot+X+d5HLUbvklYLJ9uGfcI3Y=
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U=
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/opencontainers/image-spec v1.1.0 h1:8SG7/vwALn54lVB/0yZ/MMwhFrPYtpEHQb2IpWsCzug=
@@ -94,47 +100,52 @@ github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_golang v1.19.1 h1:wZWJDwK+NameRJuPGDhlnFgx8e8HN3XHQeLaYJFJBOE=
github.com/prometheus/client_golang v1.19.1/go.mod h1:mP78NwGzrVks5S2H6ab8+ZZGJLZUq1hoULYBAYBw1Ho=
github.com/prometheus/client_golang v1.20.5 h1:cxppBPuYhUnsO6yo/aoRol4L7q7UFfdm+bR9r+8l63Y=
github.com/prometheus/client_golang v1.20.5/go.mod h1:PIEt8X02hGcP8JWbeHyeZ53Y/jReSnHgO035n//V5WE=
github.com/prometheus/client_model v0.6.1 h1:ZKSh/rekM+n3CeS952MLRAdFwIKqeY8b62p8ais2e9E=
github.com/prometheus/client_model v0.6.1/go.mod h1:OrxVMOVHjw3lKMa8+x6HeMGkHMQyHDk9E3jmP2AmGiY=
github.com/prometheus/common v0.52.2 h1:LW8Vk7BccEdONfrJBDffQGRtpSzi5CQaRZGtboOO2ck=
github.com/prometheus/common v0.52.2/go.mod h1:lrWtQx+iDfn2mbH5GUzlH9TSHyfZpHkSiG1W7y3sF2Q=
github.com/prometheus/procfs v0.13.0 h1:GqzLlQyfsPbaEHaQkO7tbDlriv/4o5Hudv6OXHGKX7o=
github.com/prometheus/procfs v0.13.0/go.mod h1:cd4PFCR54QLnGKPaKGA6l+cfuNXtht43ZKY6tow0Y1g=
github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8=
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
github.com/prometheus/common v0.59.1 h1:LXb1quJHWm1P6wq/U824uxYi4Sg0oGvNeUm1z5dJoX0=
github.com/prometheus/common v0.59.1/go.mod h1:GpWM7dewqmVYcd7SmRaiWVe9SSqjf0UrwnYnpEZNuT0=
github.com/prometheus/procfs v0.15.1 h1:YagwOFzUgYfKKHX6Dr+sHT7km/hxC76UB0learggepc=
github.com/prometheus/procfs v0.15.1/go.mod h1:fB45yRUv8NstnjriLhBQLuOUt+WW4BsoGhij/e3PBqk=
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
github.com/sammyoina/sev-snp-measure-go v0.0.0-20241107163739-38915ab517c7 h1:g+a3hLU4pl41mhP+CQu3+bDhL4HSfrnPI1BvLjKfD8Y=
github.com/sammyoina/sev-snp-measure-go v0.0.0-20241107163739-38915ab517c7/go.mod h1:dEkBe8JnxU5itNjZDEQINFd7f7l4DtjfqRuzPQcit4w=
github.com/sammyoina/sev-snp-measure-go v0.0.0-20241202151803-ef189f0ff825 h1:SqNaL9udBIc026SGNEuEuiVL0/hw9fXxM5qrFhWGkdE=
github.com/sammyoina/sev-snp-measure-go v0.0.0-20241202151803-ef189f0ff825/go.mod h1:dEkBe8JnxU5itNjZDEQINFd7f7l4DtjfqRuzPQcit4w=
github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ=
github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ=
github.com/spf13/cobra v1.8.1 h1:e5/vxKd/rZsfSJMUX1agtjeTDf+qv1/JdBF8gg5k9ZM=
github.com/spf13/cobra v1.8.1/go.mod h1:wHxEcudfqmLYa8iTfL+OuZPbBZkmvliBWKIezN3kD9Y=
github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA=
github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/spf13/pflag v1.0.6 h1:jFzHGLGAlb3ruxLB8MhbI6A8+AQX/2eW4qeyNZXNp2o=
github.com/spf13/pflag v1.0.6/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0 h1:qtFISDHKolvIxzSs0gIaiPUPR0Cucb0F2coHC7ZLdps=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0/go.mod h1:Y+Pop1Q6hCOnETWTW4NROK/q1hv50hM7yDaUTjG8lp8=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.53.0 h1:4K4tsIXefpVJtvA/8srF4V4y0akAoPHkIslgAkjixJA=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.53.0/go.mod h1:jjdQuTGVsXV4vSs+CJ2qYDeDPf9yIJV23qlIzBm73Vg=
go.opentelemetry.io/otel v1.32.0 h1:WnBN+Xjcteh0zdk01SVqV55d/m62NJLJdIyb4y/WO5U=
go.opentelemetry.io/otel v1.32.0/go.mod h1:00DCVSB0RQcnzlwyTfqtxSm+DRr9hpYrHjNGiBHVQIg=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.28.0 h1:3Q/xZUyC1BBkualc9ROb4G8qkH90LXEIICcs5zv1OYY=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.28.0/go.mod h1:s75jGIWA9OfCMzF0xr+ZgfrB5FEbbV7UuYo32ahUiFI=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.28.0 h1:j9+03ymgYhPKmeXGk5Zu+cIZOlVzd9Zv7QIiyItjFBU=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.28.0/go.mod h1:Y5+XiUG4Emn1hTfciPzGPJaSI+RpDts6BnCIir0SLqk=
go.opentelemetry.io/otel/metric v1.32.0 h1:xV2umtmNcThh2/a/aCP+h64Xx5wsj8qqnkYZktzNa0M=
go.opentelemetry.io/otel/metric v1.32.0/go.mod h1:jH7CIbbK6SH2V2wE16W05BHCtIDzauciCRLoc/SyMv8=
go.opentelemetry.io/otel/sdk v1.28.0 h1:b9d7hIry8yZsgtbmM0DKyPWMMUMlK9NEKuIG4aBqWyE=
go.opentelemetry.io/otel/sdk v1.28.0/go.mod h1:oYj7ClPUA7Iw3m+r7GeEjz0qckQRJK2B8zjcZEfu7Pg=
go.opentelemetry.io/otel/trace v1.32.0 h1:WIC9mYrXf8TmY/EXuULKc8hR17vE+Hjv2cssQDe03fM=
go.opentelemetry.io/otel/trace v1.32.0/go.mod h1:+i4rkvCraA+tG6AzwloGaCtkx53Fa+L+V8e9a7YvhT8=
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0 h1:rgMkmiGfix9vFJDcDi1PK8WEQP4FLQwLDfhp5ZLpFeE=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0/go.mod h1:ijPqXp5P6IRRByFVVg9DY8P5HkxkHE5ARIa+86aXPf4=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0 h1:DheMAlT6POBP+gh8RUH19EOTnQIor5QE0uSRPtzCpSw=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0/go.mod h1:wZcGmeVO9nzP67aYSLDqXNWK87EZWhi7JWj1v7ZXf94=
go.opentelemetry.io/otel v1.34.0 h1:zRLXxLCgL1WyKsPVrgbSdMN4c0FMkDAskSTQP+0hdUY=
go.opentelemetry.io/otel v1.34.0/go.mod h1:OWFPOQ+h4G8xpyjgqo4SxJYdDQ/qmRH+wivy7zzx9oI=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.32.0 h1:IJFEoHiytixx8cMiVAO+GmHR6Frwu+u5Ur8njpFO6Ac=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.32.0/go.mod h1:3rHrKNtLIoS0oZwkY2vxi+oJcwFRWdtUyRII+so45p8=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.32.0 h1:cMyu9O88joYEaI47CnQkxO1XZdpoTF9fEnW2duIddhw=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.32.0/go.mod h1:6Am3rn7P9TVVeXYG+wtcGE7IE1tsQ+bP3AuWcKt/gOI=
go.opentelemetry.io/otel/metric v1.34.0 h1:+eTR3U0MyfWjRDhmFMxe2SsW64QrZ84AOhvqS7Y+PoQ=
go.opentelemetry.io/otel/metric v1.34.0/go.mod h1:CEDrp0fy2D0MvkXE+dPV7cMi8tWZwX3dmaIhwPOaqHE=
go.opentelemetry.io/otel/sdk v1.32.0 h1:RNxepc9vK59A8XsgZQouW8ue8Gkb4jpWtJm9ge5lEG4=
go.opentelemetry.io/otel/sdk v1.32.0/go.mod h1:LqgegDBjKMmb2GC6/PrTnteJG39I8/vJCAP9LlJXEjU=
go.opentelemetry.io/otel/sdk/metric v1.31.0 h1:i9hxxLJF/9kkvfHppyLL55aW7iIJz4JjxTeYusH7zMc=
go.opentelemetry.io/otel/sdk/metric v1.31.0/go.mod h1:CRInTMVvNhUKgSAMbKyTMxqOBC0zgyxzW55lZzX43Y8=
go.opentelemetry.io/otel/trace v1.34.0 h1:+ouXS2V8Rd4hp4580a8q23bg0azF2nI8cqLYnC8mh/k=
go.opentelemetry.io/otel/trace v1.34.0/go.mod h1:Svm7lSjQD7kG7KJ/MUHPVXSDGz2OX4h0M2jHBhmSfRE=
go.opentelemetry.io/proto/otlp v1.3.1 h1:TrMUixzpM0yuc/znrFTP9MMRh8trP93mkCiDVeXrui0=
go.opentelemetry.io/proto/otlp v1.3.1/go.mod h1:0X1WI4de4ZsLrrJNLAQbFeLCm3T7yBkR0XqQ7niQU+8=
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
@@ -142,57 +153,53 @@ go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN8
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.29.0 h1:L5SG1JTTXupVV3n6sUqMTeWbjAyfPwoda2DLX8J8FrQ=
golang.org/x/crypto v0.29.0/go.mod h1:+F4F4N5hv6v38hfeYwTdx20oUvLLc+QfrE9Ax9HtgRg=
golang.org/x/crypto v0.32.0 h1:euUpcYgM8WcP71gNpTqQCn6rC2t6ULUPiOzfWaXVVfc=
golang.org/x/crypto v0.32.0/go.mod h1:ZnnJkOaASj8g0AjIduWNlq2NRxL0PlBrbKVyZ6V/Ugc=
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.19.0 h1:fEdghXQSo20giMthA7cd28ZC+jts4amQ3YMXiP5oMQ8=
golang.org/x/mod v0.19.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.30.0 h1:AcW1SDZMkb8IpzCdQUaIq2sP4sZ4zw+55h6ynffypl4=
golang.org/x/net v0.30.0/go.mod h1:2wGyMJ5iFasEhkwi13ChkO/t1ECNC4X4eBKkVFyYFlU=
golang.org/x/net v0.34.0 h1:Mb7Mrk043xzHgnRM88suvJFwzVrRfHEHJEl5/71CKw0=
golang.org/x/net v0.34.0/go.mod h1:di0qlW3YNM5oh6GqDGQr92MyTozJPmybPK4Ev/Gm31k=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.9.0 h1:fEo0HyrW1GIgZdpbhCRO0PkJajUS5H9IFUztCgEo2jQ=
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.10.0 h1:3NQrjDixjgGwUOCaF8w2+VYHv0Ve/vGYSbdkTa98gmQ=
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210426230700-d19ff857e887/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.27.0 h1:wBqf8DvsY9Y/2P8gAfPDEYNuS30J4lPHJxXSb/nJZ+s=
golang.org/x/sys v0.27.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/term v0.26.0 h1:WEQa6V3Gja/BhNxg540hBip/kkaYtRg3cxg4oXSw4AU=
golang.org/x/term v0.26.0/go.mod h1:Si5m1o57C5nBNQo5z1iq+XDijt21BDBDp2bK0QI8e3E=
golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/term v0.28.0 h1:/Ts8HFuMR2E6IP/jlo7QVLZHggjKQbhu/7H0LJFr3Gg=
golang.org/x/term v0.28.0/go.mod h1:Sw/lC2IAUZ92udQNf3WodGtn4k/XoLyZoh8v/8uiwek=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.20.0 h1:gK/Kv2otX8gz+wn7Rmb3vT96ZwuoxnQlY+HlJVj7Qug=
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
golang.org/x/time v0.5.0 h1:o7cqy6amK/52YcAKIPlM3a+Fpj35zvRj2TP+e1xFSfk=
golang.org/x/time v0.5.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/text v0.21.0 h1:zyQAAkrwaneQ066sspRyJaG9VNi/YJ1NfzcGB3hZ/qo=
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
golang.org/x/time v0.6.0 h1:eTDhh4ZXt5Qf0augr54TN6suAUudPcawVZeIAPU7D4U=
golang.org/x/time v0.6.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
golang.org/x/tools v0.23.0 h1:SGsXPZ+2l4JsgaCKkx+FQ9YZ5XEtA1GZYuoDjenLjvg=
golang.org/x/tools v0.23.0/go.mod h1:pnu6ufv6vQkll6szChhK3C3L/ruaIv5eBeztNG8wtsI=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/genproto/googleapis/api v0.0.0-20240903143218-8af14fe29dc1 h1:hjSy6tcFQZ171igDaN5QHOw2n6vx40juYbC/x67CEhc=
google.golang.org/genproto/googleapis/api v0.0.0-20240903143218-8af14fe29dc1/go.mod h1:qpvKtACPCQhAdu3PyQgV4l3LMXZEtft7y8QcarRsp9I=
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28 h1:XVhgTWWV3kGQlwJHR3upFWZeTsei6Oks1apkZSeonIE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28/go.mod h1:GX3210XPVPUjJbTUbvwI8f2IpZDMZuPJWDzDuebbviI=
google.golang.org/grpc v1.68.0 h1:aHQeeJbo8zAkAa3pRzrVjZlbz6uSfeOXlJNQM0RAbz0=
google.golang.org/grpc v1.68.0/go.mod h1:fmSPC5AsjSBCK54MyHRx48kpOti1/jRfOlwEWywNjWA=
google.golang.org/protobuf v1.35.2 h1:8Ar7bF+apOIoThw1EdZl0p1oWvMqTHmpA2fRTyZO8io=
google.golang.org/protobuf v1.35.2/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
google.golang.org/genproto/googleapis/api v0.0.0-20241104194629-dd2ea8efbc28 h1:M0KvPgPmDZHPlbRbaNU1APr28TvwvvdUPlSv7PUvy8g=
google.golang.org/genproto/googleapis/api v0.0.0-20241104194629-dd2ea8efbc28/go.mod h1:dguCy7UOdZhTvLzDyt15+rOrawrpM4q7DD9dQ1P11P4=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250115164207-1a7da9e5054f h1:OxYkA3wjPsZyBylwymxSHa7ViiW1Sml4ToBrncvFehI=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250115164207-1a7da9e5054f/go.mod h1:+2Yz8+CLJbIfL9z73EW45avw8Lmge3xVElCP9zEKi50=
google.golang.org/grpc v1.69.4 h1:MF5TftSMkd8GLw/m0KM6V8CMOCY6NZ1NQDPGFgbTt4A=
google.golang.org/grpc v1.69.4/go.mod h1:vyjdE6jLBI76dgpDojsFGNaHlxdjXN9ghpnd2o7JGZ4=
google.golang.org/protobuf v1.36.3 h1:82DV7MYdb8anAVi3qge1wSnMDrnKK7ebr+I0hHRN1BU=
google.golang.org/protobuf v1.36.3/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
+89
View File
@@ -0,0 +1,89 @@
#### memory config
MEMORY_SIZE=2048M
MEMORY_SLOTS=5
MAX_MEMORY=30G
#### ovmf code config
OVMF_CODE_IF=pflash
OVMF_CODE_FORMAT=raw
OVMF_CODE_UNIT=0
OVMF_CODE_FILE=/usr/share/OVMF/OVMF_CODE.fd
OVMF_CODE_READONLY=on
OVMF_VERSION=
#### ovmf vars config
OVMF_VARS_IF=pflash
OVMF_VARS_FORMAT=raw
OVMF_VARS_UNIT=1
OVMF_VARS_FILE=/usr/share/OVMF/OVMF_VARS.fd
#### net dev config
NET_DEV_ID=vmnic
NET_DEV_HOST_FWD_AGENT=7020
NET_DEV_GUEST_FWD_AGENT=7002
#### Virtio Net Pci Config
VIRTIO_NET_PCI_DISABLE_LEGACY=on
VIRTIO_NET_PCI_IOMMU_PLATFORM=true
VIRTIO_NET_PCI_ADDR=0x2
VIRTIO_NET_PCI_ROMFILE=
#### Disk image config
DISK_IMG_KERNEL_FILE=
DISK_IMG_ROOTFS_FILE=
KERNEL_COMMAND_LINE="quiet console=null"
#### Sev Config
SEV_ID=sev0
SEV_CBIT_POS=51
SEV_REDUCED_PHYS_BITS=1
SEV_HOST_DATA=
#### VSock Config
VSOCK_ID=vhost-vsock-pci0
VSOCK_GUEST_CID=3
BIN_PATH=qemu-system-x86_64
USE_SUDO=false
ENABLE_SEV=false
ENABLE_SEV_SNP=false
ENABLE_KVM=true
MACHINE=q35
CPU=EPYC
SMP_COUNT=8
SMP_MAXCPUS=64
MEM_ID=ram1
KERNEL_HASH=false
NO_GRAPHIC=true
MONITOR=pty
HOST_FWD_RANGE=6100-6200
CERTS_MOUNT=/etc/cocos/certs
ENV_MOUNT=/etc/cocos/environment
COCOS_AGENT_VERSION=v0.3.1
#### Base image URL and names
BASE_IMAGE_URL=https://cloud-images.ubuntu.com/noble/current/noble-server-cloudimg-amd64.img
BASE_IMAGE=ubuntu-base.qcow2
CUSTOM_IMAGE=ubuntu-custom.qcow2
#### Paths for OVMF firmware
OVMF_CODE=/usr/share/ovmf/x64/OVMF_CODE.4m.fd
OVMF_VARS=/usr/share/ovmf/x64/OVMF_VARS.4m.fd
#### VM parameters
VM_NAME=cocos-vm
RAM=16G
DISK_SIZE=10G # Size for root filesystem
QEMU_BINARY=qemu-system-x86_64
AGENT_GRPC_SERVER_CERT=/etc/cocos/certs/server.pem
AGENT_GRPC_SERVER_KEY=/etc/cocos/certs/key.pem
AGENT_GRPC_SERVER_CA_CERTS=/etc/cocos/ca.pem
AGENT_GRPC_CLIENT_CA_CERTS=/etc/cocos/ca.pem
+114
View File
@@ -0,0 +1,114 @@
# Agent Cloud Init Setup
## Overview
The `hal/cloud` directory contains essential files required for setting up a virtual machine (VM) with cloud-init. This setup ensures the automated installation of dependencies, configuration of the environment, and deployment of the Cocos agent as a systemd service.
### Directory Contents
- **`config.yaml`**: This YAML file provides configuration instructions for the cloud image.
- **`meta-data`**: Contains VM metadata, such as instance-specific details and identifiers.
- **`qemu.sh`**: A Bash script for downloading and configuring a cloud image, running QEMU to simulate a VM with the cloud-init configuration.
- **`.env`**: Contains environment variables for starting the VM in different modes, configuring disk space, memory allocation, and other parameters.
## Configuration
### Preparing the Cloud-Config File
The `config.yaml` file defines system configurations, including user creation, package installations, file management, and command execution.
Ensure that the cloud-config file is set up with the following configurations:
- **User Credentials**: Specify the default username and password.
- **Certificates and Keys**: Certificate files for agent for secure communication.
- **Environment Variables**: Configuration parameters required by the system.
The `config.yaml` file is divided into multiple sections, each addressing a specific aspect of the setup process.
### 1. User Configuration
This section creates a default user with specific permissions and configurations:
- Creates a user named **`cocos_user`**.
- Adds `cocos_user` to the `sudo` and `docker` groups.
- Sets a default password (should be changed for production use).
- Configures the users shell as `/bin/bash`.
### 2. Package Installation
Installs essential system packages required for various operations:
- **`curl`**: For downloading files from the web.
- **`make`**: A utility for building software.
- **`git`**: Version control system for managing code repositories.
- **`python3` and `python3-dev`**: Required for running Python-based tools.
- **`net-tools`**: Provides networking utilities such as `ifconfig` and `route`.
### 3. File Management (write_files)
Creates and configures critical files required for the setup:
- **Certificates**: Cert files (`cert.pem`, `ca.pem`, `key.pem`) located at `/etc/cocos/certs/`.
- **Environment Variables**: An env file stored at `/etc/cocos/environment`.
- **Systemd Service File**: Cocos agent service configuration file at `/etc/systemd/system/cocos-agent.service` for managing the Cocos agent.
- **Agent Scripts**:
- `agent_setup.sh`: Configures network interfaces and resizes the root filesystem.
- `agent_start_script.sh`: Sets up Docker and starts the Cocos agent.
### 4. Execution of Commands (runcmd)
A sequence of commands is executed to finalize the setup:
- Creates necessary directories: `/cocos`, `/cocos_init`, `/var/log/cocos`, `/etc/cocos`.
- Downloads and installs the Cocos agent binary.
- Installs **Wasmtime** and configures its environment variables.
- Installs **Docker** and adds `cocos_user` to the Docker group.
- Reloads systemd and enables the Cocos agent service.
## Running the Agent
To test the cloud-init configuration, execute the `qemu.sh` script to bring up a VM using QEMU:
```bash
sudo ./qemu.sh
```
**Important:** The script must be executed as root.
Once the QEMU boots the VM, the Cocos agent will run as a systemd service. The service is configured to start automatically on boot and restart in case of failure.
## Debugging and Monitoring
For troubleshooting and monitoring the Cocos agent service, use the following commands within the VM:
### Manually Start the Service
To manually start the agent service, execute:
```bash
sudo systemctl start cocos-agent.service
```
### Verify Service Status
To check if the service is running properly, use:
```bash
sudo systemctl status cocos-agent.service
```
### View Service Logs
To inspect logs generated by the agent service, execute:
```bash
journalctl -u cocos-agent.service
```
### Check Standard Output and Error Logs
To check logs stored in the system, use:
```bash
cat /var/log/cocos/agent.stdout.log
cat /var/log/cocos/agent.stderr.log
```
+174
View File
@@ -0,0 +1,174 @@
#cloud-config
package_update: true
package_upgrade: false
users:
- default
- name: cocos_user
gecos: Default User
groups:
- sudo
- docker # Add cocos user to the docker group
sudo:
- ALL=(ALL:ALL) ALL
shell: /bin/bash
chpasswd:
list: |
cocos_user:password
expire: False
ssh_pwauth: True
packages:
- curl
- make
- git
- python3
- python3-dev
- net-tools # Add net-tools to install the 'route' command
write_files:
- path: /etc/cocos/certs/cert.pem
content: |
# Add certificate content here
permissions: "0644"
- path: /etc/cocos/certs/ca.pem
content: |
# Add CA certificate content here
permissions: "0644"
- path: /etc/cocos/certs/key.pem
content: |
# Add private key content here
permissions: "0600"
- path: /etc/cocos/environment
content: |
# Add environment variables here
permissions: "0644"
- path: /etc/systemd/system/cocos-agent.service
content: |
[Unit]
Description=Cocos AI agent
After=network.target
Before=docker.service
[Service]
WorkingDirectory=/cocos
StandardOutput=file:/var/log/cocos/agent.stdout
StandardError=file:/var/log/cocos/agent.stderr
EnvironmentFile=/etc/cocos/environment
ExecStartPre=/cocos_init/agent_setup.sh
ExecStart=/cocos_init/agent_start_script.sh
Restart=always
[Install]
WantedBy=default.target
permissions: "0644"
# Agent setup script
- path: /cocos_init/agent_setup.sh
content: |
#!/bin/sh
WORK_DIR="/cocos"
# IFACES are all network interfaces excluding lo (LOOPBACK) and sit interfaces
IFACES=$(ip link show | grep -vE 'LOOPBACK|sit*' | awk -F': ' '{print $2}')
# This for loop brings up all network interfaces in IFACES and dhclient obtains an IP address for the every interface
for IFACE in $IFACES; do
STATE=$(ip link show $IFACE | grep DOWN)
if [ -n "$STATE" ]; then
ip link set $IFACE up
fi
IP_ADDR=$(ip addr show $IFACE | grep 'inet ')
if [ -z "$IP_ADDR" ]; then
dhclient $IFACE
fi
done
if [ ! -d "$WORK_DIR" ]; then
mkdir -p $WORK_DIR
fi
# Resize the root filesystem to 100% of available space
ROOT_DEV=$(findmnt / -o SOURCE -n) # Get the root filesystem device
resize2fs "$ROOT_DEV" && echo "Root filesystem resized successfully" || echo "Failed to resize root filesystem"
permissions: "0755"
# Agent start script
- path: /cocos_init/agent_start_script.sh
content: |
#!/bin/sh
# Change the docker.service file to allow Docker to run in RAM
mkdir -p /etc/systemd/system/docker.service.d
# Create or overwrite the override.conf file with the new Environment variable
tee /etc/systemd/system/docker.service.d/override.conf > /dev/null <<EOF
[Service]
Environment=DOCKER_RAMDISK=true
EOF
systemctl daemon-reload
NUM_OF_PERMITED_IFACE=1
NUM_OF_IFACE=$(ip route | grep -Eo 'dev [a-z0-9]+' | awk '{ print $2 }' | grep -v '^docker' | sort | uniq | wc -l)
if [ $NUM_OF_IFACE -gt $NUM_OF_PERMITED_IFACE ]; then
echo "More than one network interface in the VM"
exit 1
fi
DEFAULT_IFACE=$(route | grep '^default' | grep -o '[^ ]*$')
AGENT_GRPC_HOST=$(ip -4 addr show $DEFAULT_IFACE | grep inet | awk '{print $2}' | cut -d/ -f1)
export AGENT_GRPC_HOST
exec /bin/cocos-agent
permissions: "0755"
runcmd:
# Create necessary directories
- mkdir -p /cocos
- mkdir -p /cocos_init
- mkdir -p /var/log/cocos
- mkdir -p /etc/cocos
# Download the cocos-agent binary
- echo "[ COCOS AGENT SETUP ] Downloading the cocos-agent binary..."
- curl -L -O -J https://github.com/smithjilks/cocos/releases/download/v1.0.0/cocos-agent --progress-bar && echo "[ COCOS AGENT SETUP ] cocos-agent binary downloaded successfully" || echo "Failed to download cocos-agent binary"
# Install the agent binary
- echo "[ COCOS AGENT SETUP ] Installing cocos-agent binary..."
- install -D -m 0755 cocos-agent /bin/cocos-agent && echo "[ COCOS AGENT SETUP ] cocos-agent binary installed successfully" || echo "[ COCOS AGENT SETUP ] Failed to install cocos-agent binary"
# Install Wasmtime
- echo "Installing Wasmtime runtime..."
- curl https://wasmtime.dev/install.sh -sSf | bash && echo "Wasmtime installed successfully" || echo "Failed to install Wasmtime"
- echo "Configuring Wasmtime environment variables..."
- echo "export WASMTIME_HOME=$HOME/.wasmtime" >> /etc/profile.d/wasm_env.sh
- echo "export PATH=\$WASMTIME_HOME/bin:\$PATH" >> /etc/profile.d/wasm_env.sh
- . /etc/profile.d/wasm_env.sh && echo "Wasmtime environment variables configured successfully" || echo "Failed to configure Wasmtime environment variables"
# Install Docker
- echo "Starting Docker installation..."
- curl -fsSL https://get.docker.com -o get-docker.sh && echo "Docker install script downloaded successfully" || echo "Failed to download Docker install script"
- sh ./get-docker.sh && echo "Docker installed successfully" || echo "Failed to install Docker"
- usermod -aG docker cocos_user && echo "Added cocos_user to the docker group" || echo "Failed to add cocos_user to the docker group"
# Reload systemd and enable the service
- echo "[ COCOS AGENT SETUP ] Reloading systemd daemon..."
- systemctl daemon-reload && echo "[ COCOS AGENT SETUP ] Systemd daemon reloaded successfully" || echo "[ COCOS AGENT SETUP ] Failed to reload systemd daemon"
- echo "[ COCOS AGENT SETUP ] Enabling cocos-agent.service..."
- systemctl enable cocos-agent.service && echo "[ COCOS AGENT SETUP ] cocos-agent.service enabled successfully" || echo "[ COCOS AGENT SETUP ] Failed to enable cocos-agent.service"
- echo "[ COCOS AGENT SETUP ] Starting cocos-agent.service..."
- systemctl start cocos-agent.service && echo "[ COCOS AGENT SETUP ] cocos-agent.service started successfully" || echo "[ COCOS AGENT SETUP ] Failed to start cocos-agent.service"
final_message: "Cocos agent setup complete. Verify logs to confirm successful service startup."
+2
View File
@@ -0,0 +1,2 @@
instance-id: iid-cocos-vm
local-hostname: cocos-vm
+124
View File
@@ -0,0 +1,124 @@
#!/bin/bash
# Source environment variables
source ./.env
# Required commands
REQUIRED_CMDS=("wget" "cloud-localds" "$QEMU_BINARY" "qemu-img")
# Check for required commands
for cmd in "${REQUIRED_CMDS[@]}"; do
if ! command -v "$cmd" &> /dev/null; then
echo "Error: $cmd is not installed. Please install it and try again."
exit 1
fi
done
# Ensure script is run as root
if [[ $EUID -ne 0 ]]; then
echo "Error: This script must be run as root."
exit 1
fi
# Create the root filesystem image if it doesn't exist
if [ ! -f "$BASE_IMAGE" ]; then
echo "Downloading base Ubuntu image..."
wget -q "$BASE_IMAGE_URL" -O "$BASE_IMAGE" --show-progress
fi
# Create custom image
echo "Creating custom QEMU image..."
qemu-img create -f qcow2 -b "$BASE_IMAGE" -F qcow2 "$CUSTOM_IMAGE" "$DISK_SIZE"
# Cloud-init configuration files
CLOUD_CONFIG="config.yaml"
META_DATA="meta-data"
SEED_IMAGE="seed.img"
# Create seed image for cloud-init
echo "Creating seed image..."
cloud-localds "$SEED_IMAGE" "$CLOUD_CONFIG" "$META_DATA"
# Construct QEMU arguments from environment variables
construct_qemu_args() {
args=()
args+=("-name" "$VM_NAME")
# Virtualization (Enable KVM)
if [ "$ENABLE_KVM" == "true" ]; then
args+=("-enable-kvm")
fi
# Machine, CPU, RAM
if [ -n "$MACHINE" ]; then
args+=("-machine" "$MACHINE")
fi
if [ -n "$CPU" ]; then
args+=("-cpu" "$CPU")
fi
args+=("-boot" "d")
args+=("-smp" "$SMP_COUNT,maxcpus=$SMP_MAXCPUS")
args+=("-m" "$MEMORY_SIZE,slots=$MEMORY_SLOTS,maxmem=$MAX_MEMORY")
# OVMF (if applicable)
if [ "$ENABLE_SEV_SNP" != "true" ]; then
args+=("-drive" "if=$OVMF_CODE_IF,format=$OVMF_CODE_FORMAT,unit=$OVMF_CODE_UNIT,file=$OVMF_CODE,readonly=$OVMF_CODE_READONLY")
args+=("-drive" "if=$OVMF_VARS_IF,format=$OVMF_VARS_FORMAT,unit=$OVMF_VARS_UNIT,file=$OVMF_VARS")
fi
# Network configuration
args+=("-netdev" "user,id=$NET_DEV_ID,hostfwd=tcp::$NET_DEV_HOST_FWD_AGENT-:$NET_DEV_GUEST_FWD_AGENT")
args+=("-device" "virtio-net-pci,disable-legacy=$VIRTIO_NET_PCI_DISABLE_LEGACY,iommu_platform=$VIRTIO_NET_PCI_IOMMU_PLATFORM,netdev=$NET_DEV_ID,addr=$VIRTIO_NET_PCI_ADDR,romfile=$VIRTIO_NET_PCI_ROMFILE")
args+=("-device" "vhost-vsock-pci,id=$VSOCK_ID,guest-cid=$VSOCK_GUEST_CID")
# SEV (if enabled)
if [ "$ENABLE_SEV" == "true" ] || [ "$ENABLE_SEV_SNP" == "true" ]; then
sev_type="sev-guest"
kernel_hash=""
host_data=""
args+=("-machine" "confidential-guest-support=$SEV_ID,memory-backend=$MEM_ID")
if [ "$ENABLE_SEV_SNP" == "true" ]; then
args+=("-bios" "$OVMF_CODE_FILE")
sev_type="sev-snp-guest"
if [ -n "$SEV_HOST_DATA" ]; then
host_data=",host-data=$SEV_HOST_DATA"
fi
fi
if [ "$ENABLE_KERNEL_HASH" == "true" ]; then
kernel_hash=",kernel-hashes=on"
fi
args+=("-object" "memory-backend-memfd,id=$MEM_ID,size=$MEMORY_SIZE,share=true,prealloc=false")
args+=("-object" "$sev_type,id=$SEV_ID,cbitpos=$SEV_CBIT_POS,reduced-phys-bits=$SEV_REDUCED_PHYS_BITS$kernel_hash$host_data")
fi
# Disk image configuration
args+=("-drive" "file=$SEED_IMAGE,media=cdrom")
args+=("-drive" "file=$CUSTOM_IMAGE,if=none,id=disk0,format=qcow2")
args+=("-device" "virtio-scsi-pci,id=scsi,disable-legacy=on,iommu_platform=true")
args+=("-device" "scsi-hd,drive=disk0")
# Display options
if [ "$NO_GRAPHIC" == "true" ]; then
args+=("-nographic")
fi
args+=("-monitor" "$MONITOR")
args+=("-no-reboot")
args+=("-vnc" ":9")
echo "${args[@]}"
}
qemu_args=$(construct_qemu_args)
echo "Running QEMU with the following arguments: $qemu_args"
echo "Starting QEMU VM..."
$QEMU_BINARY $qemu_args
+11
View File
@@ -65,3 +65,14 @@ CONFIG_PREEMPT_DYNAMIC=n
CONFIG_DEBUG_PREEMPT=n
CONFIG_CGROUP_MISC=y
CONFIG_X86_CPUID=y
CONFIG_NET_9P=y
CONFIG_NET_9P_VIRTIO=y
CONFIG_9P_FS=y
CONFIG_9P_FS_POSIX_ACL=y
CONFIG_9P_FS_SECURITY=y
# TCG TPM
CONFIG_TCG_TPM=y
CONFIG_TCG_TPM2_HMAC=y
CONFIG_TCG_PLATFORM=y
+18 -1
View File
@@ -6,6 +6,23 @@ set -e
# Add a console on tty1
if [ -e ${TARGET_DIR}/etc/inittab ]; then
grep -qE '^tty1::' ${TARGET_DIR}/etc/inittab || \
sed -i '/GENERIC_SERIAL/a\
sed -i '/GENERIC_SERIAL/a\
tty1::respawn:/sbin/getty -L tty1 0 vt100 # QEMU graphical window' ${TARGET_DIR}/etc/inittab
fi
# Create the mount points
# Create the mount points
mkdir -p ${TARGET_DIR}/etc/certs
mkdir -p ${TARGET_DIR}/etc/cocos
# Ensure /etc/fstab exists
if [ ! -f "${TARGET_DIR}/etc/fstab" ]; then
touch "${TARGET_DIR}/etc/fstab"
fi
# Add the 9p entries to /etc/fstab
grep -q "certs_share /etc/certs" ${TARGET_DIR}/etc/fstab || \
echo "certs_share /etc/certs 9p trans=virtio,version=9p2000.L,cache=mmap 0 0" >> "${TARGET_DIR}/etc/fstab"
grep -q "env_share /etc/cocos" ${TARGET_DIR}/etc/fstab || \
echo "env_share /etc/cocos 9p trans=virtio,version=9p2000.L,cache=mmap 0 0" >> "${TARGET_DIR}/etc/fstab"
+5 -3
View File
@@ -14,6 +14,8 @@ BR2_SYSTEM_BIN_SH_BASH=y
BR2_TARGET_ROOTFS_CPIO=y
BR2_TARGET_ROOTFS_CPIO_FULL=y
BR2_TARGET_ROOTFS_CPIO_GZIP=y
BR2_TARGET_ROOTFS_OVERLAY="overlay"
BR2_PACKAGE_9PFS=y
# Image
BR2_ROOTFS_POST_BUILD_SCRIPT="$(BR2_EXTERNAL_COCOS_PATH)/board/cocos/post-build.sh"
@@ -30,9 +32,9 @@ BR2_TOOLCHAIN_HEADERS_AT_LEAST="6.12-rc6"
# Kernel
BR2_LINUX_KERNEL=y
BR2_LINUX_KERNEL_CUSTOM_GIT=y
BR2_LINUX_KERNEL_CUSTOM_REPO_URL="https://github.com/torvalds/linux.git"
BR2_LINUX_KERNEL_CUSTOM_REPO_VERSION="v6.12-rc6"
BR2_LINUX_KERNEL_VERSION="v6.12-rc6"
BR2_LINUX_KERNEL_CUSTOM_REPO_URL="https://github.com/coconut-svsm/linux.git"
BR2_LINUX_KERNEL_CUSTOM_REPO_VERSION="svsm"
BR2_LINUX_KERNEL_VERSION="svsm"
BR2_LINUX_KERNEL_PATCH=""
BR2_LINUX_KERNEL_USE_CUSTOM_CONFIG=y
BR2_LINUX_KERNEL_CUSTOM_CONFIG_FILE="$(BR2_EXTERNAL_COCOS_PATH)/board/cocos/linux.config"
+1 -2
View File
@@ -8,8 +8,7 @@ WorkingDirectory=/cocos
StandardOutput=file:/var/log/cocos/agent.stdout
StandardError=file:/var/log/cocos/agent.stderr
Environment=AGENT_GRPC_PORT=7002
Environment=AGENT_LOG_LEVEL=info
EnvironmentFile=/etc/cocos/environment
ExecStartPre=/cocos_init/agent_setup.sh
ExecStart=/cocos_init/agent_start_script.sh
+24 -5
View File
@@ -7,8 +7,9 @@ import (
"io"
"log/slog"
"github.com/ultravioletrs/cocos/agent/cvms"
"github.com/ultravioletrs/cocos/agent/events"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/types/known/timestamppb"
)
@@ -18,16 +19,17 @@ type handler struct {
opts slog.HandlerOptions
w io.Writer
cmpID string
queue chan *cvms.ClientStreamMessage
}
func NewProtoHandler(conn io.Writer, opts *slog.HandlerOptions, cmpID string) slog.Handler {
func NewProtoHandler(conn io.Writer, opts *slog.HandlerOptions, queue chan *cvms.ClientStreamMessage) slog.Handler {
if opts == nil {
opts = &slog.HandlerOptions{}
}
h := &handler{
opts: *opts,
w: conn,
cmpID: cmpID,
queue: queue,
}
return h
@@ -69,7 +71,18 @@ func (h *handler) Handle(_ context.Context, r slog.Record) error {
},
}
b, err := proto.Marshal(&agentLog)
h.queue <- &cvms.ClientStreamMessage{
Message: &cvms.ClientStreamMessage_AgentLog{
AgentLog: &cvms.AgentLog{
Timestamp: timestamp,
Message: chunk,
Level: level,
ComputationId: h.cmpID,
},
},
}
b, err := protojson.Marshal(&agentLog)
if err != nil {
return err
}
@@ -78,6 +91,11 @@ func (h *handler) Handle(_ context.Context, r slog.Record) error {
if err != nil {
return err
}
_, err = h.w.Write([]byte("\n"))
if err != nil {
return err
}
}
return nil
@@ -88,7 +106,8 @@ func (h *handler) WithAttrs(attrs []slog.Attr) slog.Handler {
}
func (h *handler) WithGroup(name string) slog.Handler {
panic("unimplemented")
h.cmpID = name
return h
}
func (h *handler) Close() error {
+6 -5
View File
@@ -11,6 +11,7 @@ import (
"github.com/absmach/magistrala/pkg/errors"
"github.com/stretchr/testify/assert"
"github.com/ultravioletrs/cocos/agent/cvms"
)
type failedWriter struct{}
@@ -21,14 +22,14 @@ func (f *failedWriter) Write(p []byte) (n int, err error) {
// TestNewProtoHandler tests the initialization of the ProtoHandler.
func TestNewProtoHandler(t *testing.T) {
handler := NewProtoHandler(io.Discard, nil, "testCmpID")
handler := NewProtoHandler(io.Discard, nil, make(chan *cvms.ClientStreamMessage))
assert.NotNil(t, handler, "Handler should not be nil")
}
// TestHandleMessageSuccess tests the handling of a message when the write succeeds.
func TestHandleMessageSuccess(t *testing.T) {
handler := NewProtoHandler(io.Discard, nil, "testCmpID")
handler := NewProtoHandler(io.Discard, nil, make(chan *cvms.ClientStreamMessage, 1))
record := slog.Record{
Time: time.Now(),
Message: "Test message",
@@ -42,7 +43,7 @@ func TestHandleMessageSuccess(t *testing.T) {
// TestHandleMessageFailure tests the caching mechanism when the write fails.
func TestHandleMessageFailure(t *testing.T) {
protohandler := NewProtoHandler(&failedWriter{}, nil, "testCmpID")
protohandler := NewProtoHandler(&failedWriter{}, nil, make(chan *cvms.ClientStreamMessage, 1))
record := slog.Record{
Time: time.Now(),
Message: "Test message",
@@ -56,7 +57,7 @@ func TestHandleMessageFailure(t *testing.T) {
// TestEnabled tests that the handler enables logging based on level.
func TestEnabled(t *testing.T) {
handler := NewProtoHandler(io.Discard, nil, "testCmpID")
handler := NewProtoHandler(io.Discard, nil, make(chan *cvms.ClientStreamMessage, 1))
assert.True(t, handler.Enabled(context.Background(), slog.LevelInfo), "Logging should be enabled for LevelInfo")
assert.False(t, handler.Enabled(context.Background(), slog.LevelDebug), "Logging should be disabled for LevelDebug by default")
@@ -66,7 +67,7 @@ func TestEnabled(t *testing.T) {
func TestCloseStopsRetry(t *testing.T) {
mockWriter := io.Discard
handler := NewProtoHandler(mockWriter, nil, "testCmpID").(*handler)
handler := NewProtoHandler(mockWriter, nil, make(chan *cvms.ClientStreamMessage, 1)).(*handler)
time.Sleep(2 * time.Second)
err := handler.Close()
-247
View File
@@ -1,247 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package vsock
import (
"context"
"encoding/binary"
"fmt"
"io"
"log"
"net"
"sync"
"sync/atomic"
"time"
"google.golang.org/protobuf/proto"
)
const (
maxRetries = 3
retryDelay = time.Second
maxMessageSize = 1 << 20 // 1 MB
ackTimeout = 5 * time.Second
maxConcurrent = 100
)
type MessageStatus int
const (
StatusPending MessageStatus = iota
StatusSent
StatusAcknowledged
StatusFailed
)
type Message struct {
ID uint32
Content []byte
Status MessageStatus
Retries int
}
type AckWriter struct {
conn net.Conn
pendingMessages chan *Message
messageStore sync.Map // map[uint32]*Message
nextID uint32
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
}
func NewAckWriter(conn net.Conn) io.WriteCloser {
ctx, cancel := context.WithCancel(context.Background())
aw := &AckWriter{
conn: conn,
pendingMessages: make(chan *Message, maxConcurrent),
nextID: 1,
ctx: ctx,
cancel: cancel,
}
aw.wg.Add(2)
go aw.sendMessages()
go aw.handleAcknowledgments()
return aw
}
func (aw *AckWriter) Write(p []byte) (int, error) {
if len(p) > maxMessageSize {
return 0, fmt.Errorf("message size exceeds maximum allowed size of %d bytes", maxMessageSize)
}
messageID := atomic.AddUint32(&aw.nextID, 1)
message := &Message{
ID: messageID,
Content: make([]byte, len(p)),
Status: StatusPending,
}
copy(message.Content, p)
aw.messageStore.Store(messageID, message)
select {
case aw.pendingMessages <- message:
return len(p), nil
case <-aw.ctx.Done():
return 0, fmt.Errorf("writer is closed")
}
}
func (aw *AckWriter) sendMessages() {
defer aw.wg.Done()
for {
select {
case <-aw.ctx.Done():
return
case msg := <-aw.pendingMessages:
if err := aw.sendWithRetry(msg); err != nil {
log.Printf("Failed to send message %d after all retries: %v", msg.ID, err)
msg.Status = StatusFailed
aw.messageStore.Store(msg.ID, msg)
}
}
}
}
func (aw *AckWriter) sendWithRetry(msg *Message) error {
for msg.Retries < maxRetries {
if err := aw.writeMessage(msg.ID, msg.Content); err != nil {
msg.Retries++
msg.Status = StatusPending
log.Printf("Error writing message %d (attempt %d): %v", msg.ID, msg.Retries, err)
time.Sleep(retryDelay)
continue
}
msg.Status = StatusSent
aw.messageStore.Store(msg.ID, msg)
return nil
}
return fmt.Errorf("max retries reached")
}
func (aw *AckWriter) writeMessage(messageID uint32, p []byte) error {
if err := binary.Write(aw.conn, binary.LittleEndian, messageID); err != nil {
return fmt.Errorf("failed to write message ID: %w", err)
}
messageLen := uint32(len(p))
if err := binary.Write(aw.conn, binary.LittleEndian, messageLen); err != nil {
return fmt.Errorf("failed to write message length: %w", err)
}
if _, err := aw.conn.Write(p); err != nil {
return fmt.Errorf("failed to write message content: %w", err)
}
return nil
}
func (aw *AckWriter) handleAcknowledgments() {
defer aw.wg.Done()
for {
select {
case <-aw.ctx.Done():
return
default:
var ackID uint32
if err := binary.Read(aw.conn, binary.LittleEndian, &ackID); err != nil {
if err == io.EOF {
log.Println("Connection closed, stopping acknowledgment handler")
return
}
log.Printf("Error reading ACK: %v", err)
time.Sleep(retryDelay)
continue
}
if msg, ok := aw.messageStore.Load(ackID); ok {
m := msg.(*Message)
m.Status = StatusAcknowledged
aw.messageStore.Store(ackID, m)
// Clean up old messages periodically
go aw.cleanupOldMessages(ackID)
} else {
log.Printf("Received ACK for unknown message ID: %d", ackID)
}
}
}
}
func (aw *AckWriter) cleanupOldMessages(currentID uint32) {
aw.messageStore.Range(func(key, value interface{}) bool {
msgID := key.(uint32)
msg := value.(*Message)
// Clean up acknowledged messages that are old
if msg.Status == StatusAcknowledged && msgID < currentID-maxConcurrent {
aw.messageStore.Delete(msgID)
}
return true
})
}
func (aw *AckWriter) Close() error {
aw.cancel()
aw.wg.Wait()
return aw.conn.Close()
}
type Reader interface {
Read() ([]byte, error)
ReadProto(msg proto.Message) error
}
type AckReader struct {
conn net.Conn
ctx context.Context
}
func NewAckReader(conn net.Conn) Reader {
return &AckReader{
conn: conn,
ctx: context.Background(),
}
}
func (ar *AckReader) ReadProto(msg proto.Message) error {
data, err := ar.Read()
if err != nil {
return fmt.Errorf("failed to read proto message: %w", err)
}
return proto.Unmarshal(data, msg)
}
func (ar *AckReader) Read() ([]byte, error) {
var messageID uint32
if err := binary.Read(ar.conn, binary.LittleEndian, &messageID); err != nil {
return nil, fmt.Errorf("error reading message ID: %w", err)
}
var messageLen uint32
if err := binary.Read(ar.conn, binary.LittleEndian, &messageLen); err != nil {
return nil, fmt.Errorf("error reading message length: %w", err)
}
if messageLen > maxMessageSize {
return nil, fmt.Errorf("message size %d exceeds maximum allowed size of %d bytes", messageLen, maxMessageSize)
}
data := make([]byte, messageLen)
if _, err := io.ReadFull(ar.conn, data); err != nil {
return nil, fmt.Errorf("error reading message content: %w", err)
}
if err := ar.sendAck(messageID); err != nil {
return nil, fmt.Errorf("error sending ACK: %w", err)
}
return data, nil
}
func (ar *AckReader) sendAck(messageID uint32) error {
return binary.Write(ar.conn, binary.LittleEndian, messageID)
}
-337
View File
@@ -1,337 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package vsock
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/ultravioletrs/cocos/manager"
"google.golang.org/protobuf/proto"
)
// MockConn implements net.Conn for testing purposes.
type MockConn struct {
ReadData []byte
WrittenData []byte
ReadErr error
WriteErr error
closed bool
mu sync.Mutex
}
func (m *MockConn) Read(b []byte) (n int, err error) {
m.mu.Lock()
defer m.mu.Unlock()
if m.closed {
return 0, io.EOF
}
if len(m.ReadData) == 0 {
return 0, io.EOF // Ensure we handle this case more predictably
}
if m.ReadErr != nil {
return 0, m.ReadErr
}
n = copy(b, m.ReadData)
m.ReadData = m.ReadData[n:]
return n, nil
}
func (m *MockConn) Write(b []byte) (n int, err error) {
m.mu.Lock()
defer m.mu.Unlock()
if m.closed {
return 0, errors.New("connection closed")
}
if m.WriteErr != nil {
return 0, m.WriteErr
}
m.WrittenData = append(m.WrittenData, b...)
return len(b), nil
}
func (m *MockConn) Close() error {
m.mu.Lock()
defer m.mu.Unlock()
m.closed = true
return nil
}
// Implement other net.Conn methods with empty implementations.
func (m *MockConn) LocalAddr() net.Addr { return nil }
func (m *MockConn) RemoteAddr() net.Addr { return nil }
func (m *MockConn) SetDeadline(t time.Time) error { return nil }
func (m *MockConn) SetReadDeadline(t time.Time) error { return nil }
func (m *MockConn) SetWriteDeadline(t time.Time) error { return nil }
func TestAckReader_Read(t *testing.T) {
tests := []struct {
name string
data []byte
wantErr bool
}{
{"Valid message", []byte("Hello, World!"), false},
{"Empty message", []byte{}, false},
{"Message at max size", make([]byte, maxMessageSize), false},
{"Message exceeds max size", make([]byte, maxMessageSize+1), true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockConn := &MockConn{}
ar := NewAckReader(mockConn)
// Prepare mock data
messageID := uint32(1)
messageLen := uint32(len(tt.data))
mockData := make([]byte, 8+len(tt.data))
binary.LittleEndian.PutUint32(mockData[:4], messageID)
binary.LittleEndian.PutUint32(mockData[4:8], messageLen)
copy(mockData[8:], tt.data)
mockConn.ReadData = mockData
data, err := ar.Read()
if (err != nil) != tt.wantErr {
t.Errorf("AckReader.Read() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !tt.wantErr {
if !bytes.Equal(data, tt.data) {
t.Errorf("AckReader.Read() got = %v, want %v", data, tt.data)
}
// Check if ACK was sent
if len(mockConn.WrittenData) != 4 {
t.Errorf("AckReader.Read() did not send ACK")
} else {
ackID := binary.LittleEndian.Uint32(mockConn.WrittenData)
if ackID != messageID {
t.Errorf("AckReader.Read() sent wrong ACK ID, got %d, want %d", ackID, messageID)
}
}
}
})
}
}
func TestAckReader_ReadProto(t *testing.T) {
tests := []struct {
name string
msg *manager.ClientStreamMessage
wantErr bool
}{
{"Valid proto message", &manager.ClientStreamMessage{}, false},
{"Empty proto message", &manager.ClientStreamMessage{}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockConn := &MockConn{}
ar := NewAckReader(mockConn)
// Prepare mock data
protoData, _ := proto.Marshal(tt.msg)
messageID := uint32(1)
messageLen := uint32(len(protoData))
mockData := make([]byte, 8+len(protoData))
binary.LittleEndian.PutUint32(mockData[:4], messageID)
binary.LittleEndian.PutUint32(mockData[4:8], messageLen)
copy(mockData[8:], protoData)
mockConn.ReadData = mockData
receivedMsg := &manager.ClientStreamMessage{}
err := ar.ReadProto(receivedMsg)
if (err != nil) != tt.wantErr {
t.Errorf("AckReader.ReadProto() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !tt.wantErr {
if receivedMsg.Message != tt.msg.Message {
t.Errorf("AckReader.ReadProto() got = %v, want %v", receivedMsg, tt.msg)
}
// Check if ACK was sent
if len(mockConn.WrittenData) != 4 {
t.Errorf("AckReader.ReadProto() did not send ACK")
} else {
ackID := binary.LittleEndian.Uint32(mockConn.WrittenData)
if ackID != messageID {
t.Errorf("AckReader.ReadProto() sent wrong ACK ID, got %d, want %d", ackID, messageID)
}
}
}
})
}
}
func TestNewAckWriter(t *testing.T) {
mockConn := &MockConn{}
writer := NewAckWriter(mockConn)
if _, ok := writer.(io.Writer); !ok {
t.Errorf("NewAckWriter() did not return an io.Writer")
}
}
func TestNewAckReader(t *testing.T) {
mockConn := &MockConn{}
reader := NewAckReader(mockConn)
assert.NotNil(t, reader)
}
func TestAckWriter_Close(t *testing.T) {
mockConn := &MockConn{}
aw := NewAckWriter(mockConn)
err := aw.Close()
if err != nil {
t.Errorf("AckWriter.Close() error = %v, wantErr %v", err, nil)
}
if !mockConn.closed {
t.Errorf("AckWriter.Close() did not close the connection")
}
}
func TestAckWriter_Write(t *testing.T) {
tests := []struct {
name string
input []byte
expectErr bool
expectedError string
}{
{
name: "Message exceeds max size",
input: make([]byte, maxMessageSize+1),
expectErr: true,
expectedError: "message size exceeds maximum allowed size",
},
{
name: "Write succeeds",
input: []byte("Hello, world!"),
expectErr: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockConn := &MockConn{
mu: sync.Mutex{},
}
writer := NewAckWriter(mockConn)
defer writer.Close()
if tt.expectErr {
writer.(*AckWriter).ctx.Done()
}
time.Sleep(100 * time.Millisecond)
n, err := writer.Write(tt.input)
if tt.expectErr {
assert.Error(t, err)
if tt.expectedError != "" {
assert.Contains(t, err.Error(), tt.expectedError)
}
assert.Zero(t, n)
} else {
assert.NoError(t, err)
assert.Equal(t, len(tt.input), n)
}
})
}
}
func TestAckWriter_CleanupOldMessages(t *testing.T) {
mockConn := &MockConn{}
writer := NewAckWriter(mockConn).(*AckWriter)
defer writer.Close()
for i := uint32(1); i <= maxConcurrent+10; i++ {
msg := &Message{
ID: i,
Content: []byte("test"),
Status: StatusAcknowledged,
}
writer.messageStore.Store(i, msg)
}
writer.cleanupOldMessages(maxConcurrent + 11)
var count int
writer.messageStore.Range(func(key, value interface{}) bool {
count++
return true
})
assert.LessOrEqual(t, count, maxConcurrent)
}
func TestAckReader_LargeMessage(t *testing.T) {
mockConn := &MockConn{}
reader := NewAckReader(mockConn)
largeMessage := make([]byte, maxMessageSize-1)
for i := range largeMessage {
largeMessage[i] = byte(i % 256)
}
messageID := uint32(1)
messageLen := uint32(len(largeMessage))
mockData := make([]byte, 8+len(largeMessage))
binary.LittleEndian.PutUint32(mockData[:4], messageID)
binary.LittleEndian.PutUint32(mockData[4:8], messageLen)
copy(mockData[8:], largeMessage)
mockConn.ReadData = mockData
data, err := reader.Read()
assert.NoError(t, err)
assert.Equal(t, largeMessage, data)
assert.Equal(t, 4, len(mockConn.WrittenData))
ackID := binary.LittleEndian.Uint32(mockConn.WrittenData)
assert.Equal(t, messageID, ackID)
}
func TestAckWriter_FailedSends(t *testing.T) {
mockConn := &MockConn{
WriteErr: errors.New("write error"),
}
writer := NewAckWriter(mockConn).(*AckWriter)
defer writer.Close()
// Add some messages to the channel
for i := 0; i < 5; i++ {
msg := &Message{
ID: uint32(i + 1),
Content: []byte(fmt.Sprintf("Message %d", i+1)),
Status: StatusPending,
}
writer.pendingMessages <- msg
}
// Wait for the messages to be sent
time.Sleep(100 * time.Millisecond)
// Check that the messages were marked as failed
writer.messageStore.Range(func(key, value interface{}) bool {
msg := value.(*Message)
assert.Equal(t, StatusFailed, msg.Status)
return true
})
}
-66
View File
@@ -1,66 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package manager
import (
"fmt"
"regexp"
"strconv"
"github.com/ultravioletrs/cocos/pkg/manager"
"google.golang.org/protobuf/types/known/timestamppb"
)
var (
errFailedToParseCID = fmt.Errorf("failed to parse computation ID")
errComputationNotFound = fmt.Errorf("computation not found")
)
func (ms *managerService) computationIDFromAddress(address string) (string, error) {
re := regexp.MustCompile(`vm\((\d+)\)`)
matches := re.FindStringSubmatch(address)
if len(matches) > 1 {
cid, err := strconv.Atoi(matches[1])
if err != nil {
return "", err
}
return ms.findComputationID(cid)
}
return "", errFailedToParseCID
}
func (ms *managerService) findComputationID(cid int) (string, error) {
ms.mu.Lock()
defer ms.mu.Unlock()
for cmpID, vm := range ms.vms {
if vm.GetCID() == cid {
return cmpID, nil
}
}
return "", errComputationNotFound
}
func (ms *managerService) reportBrokenConnection(cmpID string) {
ms.eventsChan <- &ClientStreamMessage{
Message: &ClientStreamMessage_AgentEvent{
AgentEvent: &AgentEvent{
EventType: ms.vms[cmpID].State(),
ComputationId: cmpID,
Status: manager.Disconnected.String(),
Timestamp: timestamppb.Now(),
Originator: "manager",
},
},
}
}
func (ms *managerService) ReportBrokenConnection(addr string) {
cmpID, err := ms.computationIDFromAddress(addr)
if err != nil {
ms.logger.Warn(err.Error())
return
}
ms.reportBrokenConnection(cmpID)
}
-64
View File
@@ -1,64 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package manager
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/ultravioletrs/cocos/manager/qemu"
"github.com/ultravioletrs/cocos/manager/vm"
"github.com/ultravioletrs/cocos/pkg/manager"
)
func TestComputationIDFromAddress(t *testing.T) {
ms := &managerService{
vms: map[string]vm.VM{
"comp1": qemu.NewVM(qemu.Config{VSockConfig: qemu.VSockConfig{GuestCID: 3}}, func(event interface{}) error { return nil }, "comp1"),
"comp2": qemu.NewVM(qemu.Config{VSockConfig: qemu.VSockConfig{GuestCID: 5}}, func(event interface{}) error { return nil }, "comp2"),
},
}
tests := []struct {
name string
address string
want string
wantErr bool
}{
{"Valid address", "vm(3)", "comp1", false},
{"Invalid address", "invalid", "", true},
{"Non-existent CID", "vm(10)", "", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ms.computationIDFromAddress(tt.address)
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
assert.Equal(t, tt.want, got)
}
})
}
}
func TestReportBrokenConnection(t *testing.T) {
ms := &managerService{
eventsChan: make(chan *ClientStreamMessage, 1),
vms: map[string]vm.VM{
"comp1": qemu.NewVM(qemu.Config{VSockConfig: qemu.VSockConfig{GuestCID: 3}}, func(event interface{}) error { return nil }, "comp1"),
},
}
ms.reportBrokenConnection("comp1")
select {
case msg := <-ms.eventsChan:
assert.Equal(t, "comp1", msg.GetAgentEvent().ComputationId)
assert.Equal(t, manager.Disconnected.String(), msg.GetAgentEvent().Status)
assert.Equal(t, "manager", msg.GetAgentEvent().Originator)
default:
t.Error("Expected message in eventsChan, but none received")
}
}
-242
View File
@@ -1,242 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package grpc
import (
"context"
"log/slog"
"sync"
"time"
"github.com/absmach/magistrala/pkg/errors"
"github.com/ultravioletrs/cocos/manager"
"github.com/ultravioletrs/cocos/manager/qemu"
"golang.org/x/sync/errgroup"
"google.golang.org/protobuf/proto"
)
var (
errTerminationFromServer = errors.New("server requested client termination")
errCorruptedManifest = errors.New("received manifest may be corrupted")
sendTimeout = 5 * time.Second
)
type ManagerClient struct {
stream manager.ManagerService_ProcessClient
svc manager.Service
messageQueue chan *manager.ClientStreamMessage
logger *slog.Logger
runReqManager *runRequestManager
}
// NewClient returns new gRPC client instance.
func NewClient(stream manager.ManagerService_ProcessClient, svc manager.Service, messageQueue chan *manager.ClientStreamMessage, logger *slog.Logger) ManagerClient {
return ManagerClient{
stream: stream,
svc: svc,
messageQueue: messageQueue,
logger: logger,
runReqManager: newRunRequestManager(),
}
}
func (client ManagerClient) Process(ctx context.Context, cancel context.CancelFunc) error {
eg, ctx := errgroup.WithContext(ctx)
eg.Go(func() error {
return client.handleIncomingMessages(ctx)
})
eg.Go(func() error {
return client.handleOutgoingMessages(ctx)
})
return eg.Wait()
}
func (client ManagerClient) handleIncomingMessages(ctx context.Context) error {
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
req, err := client.stream.Recv()
if err != nil {
return err
}
if err := client.processIncomingMessage(ctx, req); err != nil {
return err
}
}
}
}
func (client ManagerClient) processIncomingMessage(ctx context.Context, req *manager.ServerStreamMessage) error {
switch mes := req.Message.(type) {
case *manager.ServerStreamMessage_RunReqChunks:
return client.handleRunReqChunks(ctx, mes)
case *manager.ServerStreamMessage_TerminateReq:
return client.handleTerminateReq(mes)
case *manager.ServerStreamMessage_StopComputation:
go client.handleStopComputation(ctx, mes)
case *manager.ServerStreamMessage_AttestationPolicyReq:
go client.handleAttestationPolicyReq(ctx, mes)
case *manager.ServerStreamMessage_SvmInfoReq:
go client.handleSVMInfoReq(ctx, mes)
default:
return errors.New("unknown message type")
}
return nil
}
func (client *ManagerClient) handleRunReqChunks(ctx context.Context, mes *manager.ServerStreamMessage_RunReqChunks) error {
buffer, complete := client.runReqManager.addChunk(mes.RunReqChunks.Id, mes.RunReqChunks.Data, mes.RunReqChunks.IsLast)
if complete {
var runReq manager.ComputationRunReq
if err := proto.Unmarshal(buffer, &runReq); err != nil {
return errors.Wrap(err, errCorruptedManifest)
}
go client.executeRun(ctx, &runReq)
}
return nil
}
func (client ManagerClient) executeRun(ctx context.Context, runReq *manager.ComputationRunReq) {
port, err := client.svc.Run(ctx, runReq)
if err != nil {
client.logger.Warn(err.Error())
return
}
runRes := &manager.ClientStreamMessage_RunRes{
RunRes: &manager.RunResponse{
AgentPort: port,
ComputationId: runReq.Id,
},
}
client.sendMessage(&manager.ClientStreamMessage{Message: runRes})
}
func (client ManagerClient) handleTerminateReq(mes *manager.ServerStreamMessage_TerminateReq) error {
return errors.Wrap(errTerminationFromServer, errors.New(mes.TerminateReq.Message))
}
func (client ManagerClient) handleStopComputation(ctx context.Context, mes *manager.ServerStreamMessage_StopComputation) {
msg := &manager.ClientStreamMessage_StopComputationRes{
StopComputationRes: &manager.StopComputationResponse{
ComputationId: mes.StopComputation.ComputationId,
},
}
if err := client.svc.Stop(ctx, mes.StopComputation.ComputationId); err != nil {
msg.StopComputationRes.Message = err.Error()
}
client.sendMessage(&manager.ClientStreamMessage{Message: msg})
}
func (client ManagerClient) handleAttestationPolicyReq(ctx context.Context, mes *manager.ServerStreamMessage_AttestationPolicyReq) {
res, err := client.svc.FetchAttestationPolicy(ctx, mes.AttestationPolicyReq.Id)
if err != nil {
client.logger.Warn(err.Error())
return
}
info := &manager.ClientStreamMessage_AttestationPolicy{
AttestationPolicy: &manager.AttestationPolicy{
Info: res,
Id: mes.AttestationPolicyReq.Id,
},
}
client.sendMessage(&manager.ClientStreamMessage{Message: info})
}
func (client ManagerClient) handleSVMInfoReq(ctx context.Context, mes *manager.ServerStreamMessage_SvmInfoReq) {
ovmfVersion, cpuNum, cpuType, eosVersion := client.svc.ReturnSVMInfo(ctx)
info := &manager.ClientStreamMessage_SvmInfo{
SvmInfo: &manager.SVMInfo{
OvmfVersion: ovmfVersion,
CpuNum: int32(cpuNum),
CpuType: cpuType,
KernelCmd: qemu.KernelCommandLine,
EosVersion: eosVersion,
Id: mes.SvmInfoReq.Id,
},
}
client.sendMessage(&manager.ClientStreamMessage{Message: info})
}
func (client ManagerClient) handleOutgoingMessages(ctx context.Context) error {
for {
select {
case <-ctx.Done():
return ctx.Err()
case mes := <-client.messageQueue:
if err := client.stream.Send(mes); err != nil {
return err
}
}
}
}
func (client ManagerClient) sendMessage(mes *manager.ClientStreamMessage) {
ctx, cancel := context.WithTimeout(context.Background(), sendTimeout)
defer cancel()
select {
case client.messageQueue <- mes:
case <-ctx.Done():
client.logger.Warn("Failed to send message: timeout exceeded")
}
}
type runRequestManager struct {
requests map[string]*runRequest
mu sync.Mutex
}
type runRequest struct {
buffer []byte
lastChunk time.Time
timer *time.Timer
}
func newRunRequestManager() *runRequestManager {
return &runRequestManager{
requests: make(map[string]*runRequest),
}
}
func (m *runRequestManager) addChunk(id string, chunk []byte, isLast bool) ([]byte, bool) {
m.mu.Lock()
defer m.mu.Unlock()
req, exists := m.requests[id]
if !exists {
req = &runRequest{
buffer: make([]byte, 0),
lastChunk: time.Now(),
timer: time.AfterFunc(runReqTimeout, func() { m.timeoutRequest(id) }),
}
m.requests[id] = req
}
req.buffer = append(req.buffer, chunk...)
req.lastChunk = time.Now()
req.timer.Reset(runReqTimeout)
if isLast {
delete(m.requests, id)
req.timer.Stop()
return req.buffer, true
}
return nil, false
}
func (m *runRequestManager) timeoutRequest(id string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.requests, id)
// Log timeout or handle it as needed
}
-322
View File
@@ -1,322 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package grpc
import (
"context"
"testing"
"time"
mglog "github.com/absmach/magistrala/logger"
"github.com/absmach/magistrala/pkg/errors"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/ultravioletrs/cocos/manager"
"github.com/ultravioletrs/cocos/manager/mocks"
"github.com/ultravioletrs/cocos/manager/qemu"
"google.golang.org/grpc"
"google.golang.org/protobuf/proto"
)
type mockStream struct {
mock.Mock
grpc.ClientStream
}
func (m *mockStream) Recv() (*manager.ServerStreamMessage, error) {
args := m.Called()
return args.Get(0).(*manager.ServerStreamMessage), args.Error(1)
}
func (m *mockStream) Send(msg *manager.ClientStreamMessage) error {
args := m.Called(msg)
return args.Error(0)
}
func TestManagerClient_Process1(t *testing.T) {
tests := []struct {
name string
setupMocks func(mockStream *mockStream, mockSvc *mocks.Service)
expectError bool
errorMsg string
}{
{
name: "Stop computation",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service) {
mockStream.On("Recv").Return(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_StopComputation{
StopComputation: &manager.StopComputation{},
},
}, nil)
mockStream.On("Send", mock.Anything).Return(nil)
mockSvc.On("Stop", mock.Anything, mock.Anything).Return(nil)
},
expectError: true,
errorMsg: "context deadline exceeded",
},
{
name: "Terminate request",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service) {
mockStream.On("Recv").Return(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_TerminateReq{
TerminateReq: &manager.Terminate{},
},
}, nil)
},
expectError: true,
errorMsg: errTerminationFromServer.Error(),
},
{
name: "Attestation Policy request",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service) {
mockStream.On("Recv").Return(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_AttestationPolicyReq{
AttestationPolicyReq: &manager.AttestationPolicyReq{},
},
}, nil)
mockStream.On("Send", mock.Anything).Return(nil).Once()
mockSvc.On("FetchAttestationPolicy", mock.Anything, mock.Anything).Return(nil, assert.AnError)
},
expectError: true,
},
{
name: "Run request chunks",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service) {
mockStream.On("Recv").Return(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_RunReqChunks{
RunReqChunks: &manager.RunReqChunks{},
},
}, nil)
mockStream.On("Send", mock.Anything).Return(nil).Once()
mockSvc.On("Run", mock.Anything, mock.Anything).Return("", assert.AnError).Once()
},
expectError: true,
},
{
name: "Receive error",
setupMocks: func(mockStream *mockStream, mockSvc *mocks.Service) {
mockStream.On("Recv").Return(&manager.ServerStreamMessage{}, assert.AnError)
},
expectError: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
tc.setupMocks(mockStream, mockSvc)
err := client.Process(ctx, cancel)
if tc.expectError {
assert.Error(t, err)
if tc.errorMsg != "" {
assert.Contains(t, err.Error(), tc.errorMsg)
}
} else {
assert.NoError(t, err)
}
})
}
}
func TestManagerClient_handleRunReqChunks(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
runReq := &manager.ComputationRunReq{
Id: "test-id",
}
runReqBytes, _ := proto.Marshal(runReq)
chunk1 := &manager.ServerStreamMessage_RunReqChunks{
RunReqChunks: &manager.RunReqChunks{
Id: "chunk-1",
Data: runReqBytes[:len(runReqBytes)/2],
IsLast: false,
},
}
chunk2 := &manager.ServerStreamMessage_RunReqChunks{
RunReqChunks: &manager.RunReqChunks{
Id: "chunk-1",
Data: runReqBytes[len(runReqBytes)/2:],
IsLast: true,
},
}
mockSvc.On("Run", mock.Anything, mock.AnythingOfType("*manager.ComputationRunReq")).Return("8080", nil)
err := client.handleRunReqChunks(context.Background(), chunk1)
assert.NoError(t, err)
err = client.handleRunReqChunks(context.Background(), chunk2)
assert.NoError(t, err)
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
runRes, ok := msg.Message.(*manager.ClientStreamMessage_RunRes)
assert.True(t, ok)
assert.Equal(t, "8080", runRes.RunRes.AgentPort)
assert.Equal(t, "test-id", runRes.RunRes.ComputationId)
}
func TestManagerClient_handleTerminateReq(t *testing.T) {
client := ManagerClient{}
terminateReq := &manager.ServerStreamMessage_TerminateReq{
TerminateReq: &manager.Terminate{
Message: "Test termination",
},
}
err := client.handleTerminateReq(terminateReq)
assert.Error(t, err)
assert.Contains(t, err.Error(), "Test termination")
assert.True(t, errors.Contains(err, errTerminationFromServer))
}
func TestManagerClient_handleStopComputation(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
stopReq := &manager.ServerStreamMessage_StopComputation{
StopComputation: &manager.StopComputation{
ComputationId: "test-comp-id",
},
}
mockSvc.On("Stop", mock.Anything, "test-comp-id").Return(nil)
client.handleStopComputation(context.Background(), stopReq)
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
stopRes, ok := msg.Message.(*manager.ClientStreamMessage_StopComputationRes)
assert.True(t, ok)
assert.Equal(t, "test-comp-id", stopRes.StopComputationRes.ComputationId)
assert.Empty(t, stopRes.StopComputationRes.Message)
}
func TestManagerClient_handleAttestationPolicyReq(t *testing.T) {
t.Run("success", func(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
infoReq := &manager.ServerStreamMessage_AttestationPolicyReq{
AttestationPolicyReq: &manager.AttestationPolicyReq{
Id: "test-info-id",
},
}
mockSvc.On("FetchAttestationPolicy", context.Background(), infoReq.AttestationPolicyReq.Id).Return([]byte("test-attestation-policy"), nil)
client.handleAttestationPolicyReq(context.Background(), infoReq)
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
infoRes, ok := msg.Message.(*manager.ClientStreamMessage_AttestationPolicy)
assert.True(t, ok)
assert.Equal(t, "test-info-id", infoRes.AttestationPolicy.Id)
assert.Equal(t, []byte("test-attestation-policy"), infoRes.AttestationPolicy.Info)
})
t.Run("error", func(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
infoReq := &manager.ServerStreamMessage_AttestationPolicyReq{
AttestationPolicyReq: &manager.AttestationPolicyReq{
Id: "test-info-id",
},
}
mockSvc.On("FetchAttestationPolicy", context.Background(), infoReq.AttestationPolicyReq.Id).Return(nil, assert.AnError)
client.handleAttestationPolicyReq(context.Background(), infoReq)
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 0)
})
}
func TestManagerClient_handleSVMInfoReq(t *testing.T) {
mockStream := new(mockStream)
mockSvc := new(mocks.Service)
messageQueue := make(chan *manager.ClientStreamMessage, 10)
logger := mglog.NewMock()
client := NewClient(mockStream, mockSvc, messageQueue, logger)
mockSvc.On("ReturnSVMInfo", context.Background()).Return("edk2-stable202408", 4, "EPYC", "")
client.handleSVMInfoReq(context.Background(), &manager.ServerStreamMessage_SvmInfoReq{SvmInfoReq: &manager.SVMInfoReq{Id: "test-svm-info-id"}})
// Wait for the goroutine to finish
time.Sleep(50 * time.Millisecond)
mockSvc.AssertExpectations(t)
assert.Len(t, messageQueue, 1)
msg := <-messageQueue
infoRes, ok := msg.Message.(*manager.ClientStreamMessage_SvmInfo)
assert.True(t, ok)
assert.Equal(t, "edk2-stable202408", infoRes.SvmInfo.OvmfVersion)
assert.Equal(t, int32(4), infoRes.SvmInfo.CpuNum)
assert.Equal(t, "EPYC", infoRes.SvmInfo.CpuType)
assert.Equal(t, "", infoRes.SvmInfo.EosVersion)
assert.Equal(t, qemu.KernelCommandLine, infoRes.SvmInfo.KernelCmd)
}
func TestManagerClient_timeoutRequest(t *testing.T) {
rm := newRunRequestManager()
rm.requests["test-id"] = &runRequest{
timer: time.NewTimer(100 * time.Millisecond),
buffer: []byte("test-data"),
lastChunk: time.Now(),
}
rm.timeoutRequest("test-id")
assert.Len(t, rm.requests, 0)
}
+43 -104
View File
@@ -3,17 +3,11 @@
package grpc
import (
"bytes"
"context"
"errors"
"io"
"time"
"github.com/ultravioletrs/cocos/manager"
"golang.org/x/sync/errgroup"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/emptypb"
)
var (
@@ -21,113 +15,58 @@ var (
ErrUnexpectedMsg = errors.New("unknown message type")
)
const (
bufferSize = 1024 * 1024 // 1 MB
runReqTimeout = 30 * time.Second
)
type SendFunc func(*manager.ServerStreamMessage) error
type grpcServer struct {
manager.UnimplementedManagerServiceServer
incoming chan *manager.ClientStreamMessage
svc Service
}
type Service interface {
Run(ctx context.Context, ipAddress string, sendMessage SendFunc, authInfo credentials.AuthInfo)
svc manager.Service
}
// NewServer returns new AuthServiceServer instance.
func NewServer(incoming chan *manager.ClientStreamMessage, svc Service) manager.ManagerServiceServer {
func NewServer(svc manager.Service) manager.ManagerServiceServer {
return &grpcServer{
incoming: incoming,
svc: svc,
svc: svc,
}
}
func (s *grpcServer) Process(stream manager.ManagerService_ProcessServer) error {
client, ok := peer.FromContext(stream.Context())
if !ok {
return errors.New("failed to get peer info")
}
eg, ctx := errgroup.WithContext(stream.Context())
eg.Go(func() error {
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
req, err := stream.Recv()
if err != nil {
return err
}
s.incoming <- req
}
}
})
eg.Go(func() error {
sendMessage := func(msg *manager.ServerStreamMessage) error {
select {
case <-ctx.Done():
return ctx.Err()
default:
switch m := msg.Message.(type) {
case *manager.ServerStreamMessage_RunReq:
return s.sendRunReqInChunks(stream, m.RunReq)
default:
return stream.Send(msg)
}
}
}
s.svc.Run(ctx, client.Addr.String(), sendMessage, client.AuthInfo)
return nil
})
return eg.Wait()
}
func (s *grpcServer) sendRunReqInChunks(stream manager.ManagerService_ProcessServer, runReq *manager.ComputationRunReq) error {
data, err := proto.Marshal(runReq)
func (s *grpcServer) CreateVm(ctx context.Context, req *manager.CreateReq) (*manager.CreateRes, error) {
port, id, err := s.svc.CreateVM(ctx, req)
if err != nil {
return err
return nil, err
}
dataBuffer := bytes.NewBuffer(data)
buf := make([]byte, bufferSize)
for {
n, err := dataBuffer.Read(buf)
isLast := false
if err == io.EOF {
isLast = true
} else if err != nil {
return err
}
chunk := &manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_RunReqChunks{
RunReqChunks: &manager.RunReqChunks{
Id: runReq.Id,
Data: buf[:n],
IsLast: isLast,
},
},
}
if err := stream.Send(chunk); err != nil {
return err
}
if isLast {
break
}
}
return nil
return &manager.CreateRes{
ForwardedPort: port,
SvmId: id,
}, nil
}
func (s *grpcServer) RemoveVm(ctx context.Context, req *manager.RemoveReq) (*emptypb.Empty, error) {
if err := s.svc.RemoveVM(ctx, req.SvmId); err != nil {
return nil, err
}
return &emptypb.Empty{}, nil
}
func (s *grpcServer) SVMInfo(ctx context.Context, req *manager.SVMInfoReq) (*manager.SVMInfoRes, error) {
ovmf, cpunum, cputype, eosversion := s.svc.ReturnSVMInfo(ctx)
return &manager.SVMInfoRes{
OvmfVersion: ovmf,
CpuNum: int32(cpunum),
CpuType: cputype,
EosVersion: eosversion,
Id: req.Id,
}, nil
}
func (s *grpcServer) AttestationPolicy(ctx context.Context, req *manager.AttestationPolicyReq) (*manager.AttestationPolicyRes, error) {
policy, err := s.svc.FetchAttestationPolicy(ctx, req.Id)
if err != nil {
return nil, err
}
return &manager.AttestationPolicyRes{
Info: policy,
Id: req.Id,
}, nil
}
+6 -10
View File
@@ -27,9 +27,9 @@ func LoggingMiddleware(svc manager.Service, logger *slog.Logger) manager.Service
return &loggingMiddleware{logger, svc}
}
func (lm *loggingMiddleware) Run(ctx context.Context, mc *manager.ComputationRunReq) (agentAddr string, err error) {
func (lm *loggingMiddleware) CreateVM(ctx context.Context, req *manager.CreateReq) (agentAddr string, id string, err error) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method Run for computation took %s to complete", time.Since(begin))
message := fmt.Sprintf("Method CreateVM for id %s on port %s took %s to complete", id, agentAddr, time.Since(begin))
if err != nil {
lm.logger.Warn(fmt.Sprintf("%s with error: %s.", message, err))
return
@@ -37,12 +37,12 @@ func (lm *loggingMiddleware) Run(ctx context.Context, mc *manager.ComputationRun
lm.logger.Info(message)
}(time.Now())
return lm.svc.Run(ctx, mc)
return lm.svc.CreateVM(ctx, req)
}
func (lm *loggingMiddleware) Stop(ctx context.Context, computationID string) (err error) {
func (lm *loggingMiddleware) RemoveVM(ctx context.Context, id string) (err error) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method Stop for computation took %s to complete", time.Since(begin))
message := fmt.Sprintf("Method RemoveVM for vm %s took %s to complete", id, time.Since(begin))
if err != nil {
lm.logger.Warn(fmt.Sprintf("%s with error: %s.", message, err))
return
@@ -50,7 +50,7 @@ func (lm *loggingMiddleware) Stop(ctx context.Context, computationID string) (er
lm.logger.Info(message)
}(time.Now())
return lm.svc.Stop(ctx, computationID)
return lm.svc.RemoveVM(ctx, id)
}
func (lm *loggingMiddleware) FetchAttestationPolicy(ctx context.Context, cmpId string) (body []byte, err error) {
@@ -67,10 +67,6 @@ func (lm *loggingMiddleware) FetchAttestationPolicy(ctx context.Context, cmpId s
return lm.svc.FetchAttestationPolicy(ctx, cmpId)
}
func (lm *loggingMiddleware) ReportBrokenConnection(addr string) {
lm.svc.ReportBrokenConnection(addr)
}
func (lm *loggingMiddleware) ReturnSVMInfo(ctx context.Context) (string, int, string, string) {
defer func(begin time.Time) {
message := fmt.Sprintf("Method ReturnSVMInfo for computation took %s to complete", time.Since(begin))
+4 -8
View File
@@ -32,22 +32,22 @@ func MetricsMiddleware(svc manager.Service, counter metrics.Counter, latency met
}
}
func (ms *metricsMiddleware) Run(ctx context.Context, mc *manager.ComputationRunReq) (string, error) {
func (ms *metricsMiddleware) CreateVM(ctx context.Context, req *manager.CreateReq) (string, string, error) {
defer func(begin time.Time) {
ms.counter.With("method", "Run").Add(1)
ms.latency.With("method", "Run").Observe(time.Since(begin).Seconds())
}(time.Now())
return ms.svc.Run(ctx, mc)
return ms.svc.CreateVM(ctx, req)
}
func (ms *metricsMiddleware) Stop(ctx context.Context, computationID string) error {
func (ms *metricsMiddleware) RemoveVM(ctx context.Context, computationID string) error {
defer func(begin time.Time) {
ms.counter.With("method", "Stop").Add(1)
ms.latency.With("method", "Stop").Observe(time.Since(begin).Seconds())
}(time.Now())
return ms.svc.Stop(ctx, computationID)
return ms.svc.RemoveVM(ctx, computationID)
}
func (ms *metricsMiddleware) FetchAttestationPolicy(ctx context.Context, cmpId string) ([]byte, error) {
@@ -59,10 +59,6 @@ func (ms *metricsMiddleware) FetchAttestationPolicy(ctx context.Context, cmpId s
return ms.svc.FetchAttestationPolicy(ctx, cmpId)
}
func (ms *metricsMiddleware) ReportBrokenConnection(addr string) {
ms.svc.ReportBrokenConnection(addr)
}
func (ms *metricsMiddleware) ReturnSVMInfo(ctx context.Context) (string, int, string, string) {
defer func(begin time.Time) {
ms.counter.With("method", "ReturnSVMInfo").Add(1)
+14 -8
View File
@@ -34,17 +34,21 @@ func (ms *managerService) FetchAttestationPolicy(_ context.Context, computationI
return nil, fmt.Errorf("computationId %s not found", computationId)
}
config, ok := vm.GetConfig().(qemu.Config)
vmi, ok := vm.GetConfig().(qemu.VMInfo)
if !ok {
return nil, fmt.Errorf("failed to cast config to qemu.Config")
return nil, fmt.Errorf("failed to cast config to qemu.VMInfo")
}
ms.ap.Lock()
_, err := cmd.Output()
ms.ap.Unlock()
if err != nil {
return nil, err
}
ms.ap.Lock()
f, err := os.ReadFile("./attestation_policy.json")
ms.ap.Unlock()
if err != nil {
return nil, err
}
@@ -57,13 +61,13 @@ func (ms *managerService) FetchAttestationPolicy(_ context.Context, computationI
var measurement []byte
switch {
case config.EnableSEV:
measurement, err = guest.CalcLaunchDigest(guest.SEV, config.SMPCount, uint64(cpuid.CpuSigs[ms.qemuCfg.CPU]), config.OVMFCodeConfig.File, config.KernelFile, config.RootFsFile, strconv.Quote(qemu.KernelCommandLine), defGuestFeatures, "", vmmtypes.QEMU, false, "", 0)
case vmi.Config.EnableSEV:
measurement, err = guest.CalcLaunchDigest(guest.SEV, vmi.Config.SMPCount, uint64(cpuid.CpuSigs[ms.qemuCfg.CPU]), vmi.Config.OVMFCodeConfig.File, vmi.Config.KernelFile, vmi.Config.RootFsFile, strconv.Quote(qemu.KernelCommandLine), defGuestFeatures, "", vmmtypes.QEMU, false, "", 0)
if err != nil {
return nil, err
}
case config.EnableSEVSNP:
measurement, err = guest.CalcLaunchDigest(guest.SEV_SNP, config.SMPCount, uint64(cpuid.CpuSigs[config.CPU]), config.OVMFCodeConfig.File, config.KernelFile, config.RootFsFile, strconv.Quote(qemu.KernelCommandLine), defGuestFeatures, "", vmmtypes.QEMU, false, "", 0)
case vmi.Config.EnableSEVSNP:
measurement, err = guest.CalcLaunchDigest(guest.SEV_SNP, vmi.Config.SMPCount, uint64(cpuid.CpuSigs[vmi.Config.CPU]), vmi.Config.OVMFCodeConfig.File, vmi.Config.KernelFile, vmi.Config.RootFsFile, strconv.Quote(qemu.KernelCommandLine), defGuestFeatures, "", vmmtypes.QEMU, false, "", 0)
if err != nil {
return nil, err
}
@@ -72,14 +76,16 @@ func (ms *managerService) FetchAttestationPolicy(_ context.Context, computationI
attestationPolicy.Policy.Measurement = measurement
}
if config.HostData != "" {
hostData, err := base64.StdEncoding.DecodeString(config.HostData)
if vmi.Config.HostData != "" {
hostData, err := base64.StdEncoding.DecodeString(vmi.Config.HostData)
if err != nil {
return nil, err
}
attestationPolicy.Policy.HostData = hostData
}
attestationPolicy.Policy.MinimumLaunchTcb = vmi.LaunchTCB
f, err = protojson.Marshal(&attestationPolicy)
if err != nil {
return nil, err
+32 -20
View File
@@ -15,7 +15,7 @@ import (
"github.com/ultravioletrs/cocos/manager/vm/mocks"
)
func createDummyAttestationPolicyBinary(t *testing.T, behavior string) string {
func CreateDummyAttestationPolicyBinary(t *testing.T, behavior string) string {
var content []byte
switch behavior {
case "success":
@@ -55,13 +55,16 @@ func TestFetchAttestationPolicy(t *testing.T) {
name: "Valid SEV configuration",
computationId: "sev-computation",
binaryBehavior: "success",
vmConfig: qemu.Config{
EnableSEV: true,
SMPCount: 2,
CPU: "EPYC",
OVMFCodeConfig: qemu.OVMFCodeConfig{
File: "/path/to/OVMF_CODE.fd",
vmConfig: qemu.VMInfo{
Config: qemu.Config{
EnableSEV: true,
SMPCount: 2,
CPU: "EPYC",
OVMFCodeConfig: qemu.OVMFCodeConfig{
File: "/path/to/OVMF_CODE.fd",
},
},
LaunchTCB: 0,
},
expectedError: "open /path/to/OVMF_CODE.fd: no such file or directory",
},
@@ -69,13 +72,16 @@ func TestFetchAttestationPolicy(t *testing.T) {
name: "Valid SEV-SNP configuration",
computationId: "sev-snp-computation",
binaryBehavior: "success",
vmConfig: qemu.Config{
EnableSEVSNP: true,
SMPCount: 4,
CPU: "EPYC-v2",
OVMFCodeConfig: qemu.OVMFCodeConfig{
File: "/path/to/OVMF_CODE_SNP.fd",
vmConfig: qemu.VMInfo{
Config: qemu.Config{
EnableSEVSNP: true,
SMPCount: 4,
CPU: "EPYC-v2",
OVMFCodeConfig: qemu.OVMFCodeConfig{
File: "/path/to/OVMF_CODE_SNP.fd",
},
},
LaunchTCB: 0,
},
expectedError: "open /path/to/OVMF_CODE_SNP.fd: no such file or director",
},
@@ -83,7 +89,7 @@ func TestFetchAttestationPolicy(t *testing.T) {
name: "Invalid computation ID",
computationId: "non-existent",
binaryBehavior: "success",
vmConfig: qemu.Config{},
vmConfig: qemu.VMInfo{Config: qemu.Config{}, LaunchTCB: 0},
expectedError: "computationId non-existent not found",
},
{
@@ -91,14 +97,17 @@ func TestFetchAttestationPolicy(t *testing.T) {
computationId: "invalid-config",
binaryBehavior: "success",
vmConfig: struct{}{},
expectedError: "failed to cast config to qemu.Config",
expectedError: "failed to cast config to qemu.VMInfo",
},
{
name: "Binary execution failure",
computationId: "binary-fail",
binaryBehavior: "fail",
vmConfig: qemu.Config{
EnableSEV: true,
vmConfig: qemu.VMInfo{
Config: qemu.Config{
EnableSEV: true,
},
LaunchTCB: 0,
},
expectedError: "exit status 1",
},
@@ -106,8 +115,11 @@ func TestFetchAttestationPolicy(t *testing.T) {
name: "JSON file not created",
computationId: "no-json",
binaryBehavior: "no_json",
vmConfig: qemu.Config{
EnableSEV: true,
vmConfig: qemu.VMInfo{
Config: qemu.Config{
EnableSEV: true,
},
LaunchTCB: 0,
},
expectedError: "no such file or directory",
},
@@ -115,7 +127,7 @@ func TestFetchAttestationPolicy(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
tempDir := createDummyAttestationPolicyBinary(t, tc.binaryBehavior)
tempDir := CreateDummyAttestationPolicyBinary(t, tc.binaryBehavior)
defer os.RemoveAll(tempDir)
ms := &managerService{
-9
View File
@@ -1,9 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package events
import "context"
type Listener interface {
Listen(ctx context.Context)
}
-125
View File
@@ -1,125 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package events
import (
"context"
"log/slog"
"net"
"github.com/mdlayher/vsock"
agentevents "github.com/ultravioletrs/cocos/agent/events"
internalvsock "github.com/ultravioletrs/cocos/internal/vsock"
"github.com/ultravioletrs/cocos/manager"
"google.golang.org/protobuf/proto"
)
const ManagerVsockPort = 9997
type ReportBrokenConnectionFunc func(address string)
type events struct {
reportBrokenConnection ReportBrokenConnectionFunc
lis net.Listener
logger *slog.Logger
eventsChan chan *manager.ClientStreamMessage
}
func New(logger *slog.Logger, reportBrokenConnection ReportBrokenConnectionFunc, eventsChan chan *manager.ClientStreamMessage) (Listener, error) {
l, err := vsock.Listen(ManagerVsockPort, nil)
if err != nil {
return nil, err
}
return &events{
lis: l,
reportBrokenConnection: reportBrokenConnection,
logger: logger,
eventsChan: eventsChan,
}, nil
}
func (e *events) Listen(ctx context.Context) {
for {
select {
case <-ctx.Done():
e.logger.Info("Listener shutting down")
return
default:
conn, err := e.lis.Accept()
if err != nil {
e.logger.Warn(err.Error())
continue
}
go e.handleConnection(conn)
}
}
}
func (e *events) handleConnection(conn net.Conn) {
defer conn.Close()
ackReader := internalvsock.NewAckReader(conn)
for {
var message agentevents.EventsLogs
data, err := ackReader.Read()
if err != nil {
go e.reportBrokenConnection(conn.RemoteAddr().String())
e.logger.Warn(err.Error())
return
}
if err := proto.Unmarshal(data, &message); err != nil {
e.logger.Warn(err.Error())
continue
}
var mes manager.ClientStreamMessage
args := []any{}
switch message.Message.(type) {
case *agentevents.EventsLogs_AgentEvent:
args = append(args, slog.Group("agent-event",
slog.String("event-type", message.GetAgentEvent().GetEventType()),
slog.String("computation-id", message.GetAgentEvent().GetComputationId()),
slog.String("status", message.GetAgentEvent().GetStatus()),
slog.String("originator", message.GetAgentEvent().GetOriginator()),
slog.String("timestamp", message.GetAgentEvent().GetTimestamp().String()),
slog.String("details", string(message.GetAgentEvent().GetDetails()))))
mes = manager.ClientStreamMessage{
Message: &manager.ClientStreamMessage_AgentEvent{
AgentEvent: &manager.AgentEvent{
EventType: message.GetAgentEvent().GetEventType(),
ComputationId: message.GetAgentEvent().GetComputationId(),
Status: message.GetAgentEvent().GetStatus(),
Originator: message.GetAgentEvent().GetOriginator(),
Timestamp: message.GetAgentEvent().GetTimestamp(),
Details: message.GetAgentEvent().GetDetails(),
},
},
}
case *agentevents.EventsLogs_AgentLog:
args = append(args, slog.Group("agent-log",
slog.String("computation-id", message.GetAgentLog().GetComputationId()),
slog.String("level", message.GetAgentLog().GetLevel()),
slog.String("timestamp", message.GetAgentLog().GetTimestamp().String()),
slog.String("message", message.GetAgentLog().GetMessage())))
mes = manager.ClientStreamMessage{
Message: &manager.ClientStreamMessage_AgentLog{
AgentLog: &manager.AgentLog{
ComputationId: message.GetAgentLog().GetComputationId(),
Level: message.GetAgentLog().GetLevel(),
Timestamp: message.GetAgentLog().GetTimestamp(),
Message: message.GetAgentLog().GetMessage(),
},
},
}
}
e.eventsChan <- &mes
e.logger.Info("", args...)
}
}
-295
View File
@@ -1,295 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package events
import (
"bytes"
"context"
"encoding/binary"
"fmt"
"log/slog"
"net"
"os"
"testing"
"time"
mglog "github.com/absmach/magistrala/logger"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/ultravioletrs/cocos/manager"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/timestamppb"
)
type MockVsockListener struct {
mock.Mock
}
func (m *MockVsockListener) Accept() (net.Conn, error) {
args := m.Called()
return args.Get(0).(net.Conn), args.Error(1)
}
func (m *MockVsockListener) Close() error {
args := m.Called()
return args.Error(0)
}
func (m *MockVsockListener) Addr() net.Addr {
args := m.Called()
return args.Get(0).(net.Addr)
}
var _ net.Conn = (*MockConn)(nil)
type MockConn struct {
mock.Mock
}
func (m *MockConn) Read(b []byte) (n int, err error) {
args := m.Called(b)
return args.Int(0), args.Error(1)
}
func (m *MockConn) Write(b []byte) (n int, err error) {
args := m.Called(b)
return args.Int(0), args.Error(1)
}
func (m *MockConn) Close() error {
args := m.Called()
return args.Error(0)
}
func (m *MockConn) LocalAddr() net.Addr {
args := m.Called()
return args.Get(0).(net.Addr)
}
func (m *MockConn) RemoteAddr() net.Addr {
args := m.Called()
return args.Get(0).(net.Addr)
}
func (m *MockConn) SetDeadline(t time.Time) error {
args := m.Called(t)
return args.Error(0)
}
func (m *MockConn) SetReadDeadline(t time.Time) error {
args := m.Called(t)
return args.Error(0)
}
func (m *MockConn) SetWriteDeadline(t time.Time) error {
args := m.Called(t)
return args.Error(0)
}
func TestNew(t *testing.T) {
logger := &slog.Logger{}
reportBrokenConnection := func(address string) {}
eventsChan := make(chan *manager.ClientStreamMessage)
e, err := New(logger, reportBrokenConnection, eventsChan)
if vsockDeviceExists() {
assert.NoError(t, err)
assert.NotNil(t, e)
assert.IsType(t, &events{}, e)
} else {
assert.Error(t, err)
}
}
func TestListen(t *testing.T) {
mockListener := new(MockVsockListener)
mockConn := new(MockConn)
e := &events{
lis: mockListener,
logger: mglog.NewMock(),
}
mockListener.On("Accept").Return(mockConn, fmt.Errorf("mock error")).Once()
mockListener.On("Accept").Return(mockConn, nil)
mockConn.On("Close").Return(nil)
mockConn.On("Read", mock.Anything).Return(0, nil)
go e.Listen(context.Background())
time.Sleep(100 * time.Millisecond)
mockListener.AssertExpectations(t)
}
func TestListenContextDone(t *testing.T) {
mockListener := new(MockVsockListener)
mockConn := new(MockConn)
e := &events{
lis: mockListener,
logger: mglog.NewMock(),
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
mockListener.On("Accept").Return(mockConn, nil)
e.Listen(ctx)
time.Sleep(100 * time.Millisecond)
}
func vsockDeviceExists() bool {
fs, err := os.Stat("/dev/vsock")
if err != nil {
return false
}
if fs.Mode()&os.ModeDevice == 0 {
return false
}
return true
}
type MockConnWithBuffer struct {
mock.Mock
readBuf *bytes.Buffer
writeBuf *bytes.Buffer
}
func NewMockConnWithBuffer() *MockConnWithBuffer {
return &MockConnWithBuffer{
readBuf: new(bytes.Buffer),
writeBuf: new(bytes.Buffer),
}
}
func (m *MockConnWithBuffer) Read(b []byte) (n int, err error) {
return m.readBuf.Read(b)
}
func (m *MockConnWithBuffer) Write(b []byte) (n int, err error) {
return m.writeBuf.Write(b)
}
func (m *MockConnWithBuffer) Close() error {
return nil
}
func (m *MockConnWithBuffer) LocalAddr() net.Addr {
return nil
}
func (m *MockConnWithBuffer) RemoteAddr() net.Addr {
return &net.IPAddr{IP: net.ParseIP("localhost")}
}
func (m *MockConnWithBuffer) SetDeadline(t time.Time) error {
return nil
}
func (m *MockConnWithBuffer) SetReadDeadline(t time.Time) error {
return nil
}
func (m *MockConnWithBuffer) SetWriteDeadline(t time.Time) error {
return nil
}
func TestHandleConnection(t *testing.T) {
tests := []struct {
name string
message *manager.ClientStreamMessage
}{
{
name: "handle agent event",
message: &manager.ClientStreamMessage{
Message: &manager.ClientStreamMessage_AgentEvent{
AgentEvent: &manager.AgentEvent{
EventType: "test_event",
ComputationId: "test_computation",
Status: "test_status",
Originator: "test_originator",
Timestamp: timestamppb.Now(),
Details: []byte("test_details"),
},
},
},
},
{
name: "handle agent log",
message: &manager.ClientStreamMessage{
Message: &manager.ClientStreamMessage_AgentLog{
AgentLog: &manager.AgentLog{
ComputationId: "test_computation",
Timestamp: timestamppb.Now(),
},
},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockConn := NewMockConnWithBuffer()
eventsChan := make(chan *manager.ClientStreamMessage, 1)
e := &events{
logger: mglog.NewMock(),
eventsChan: eventsChan,
reportBrokenConnection: func(address string) {},
}
data, err := proto.Marshal(tt.message)
assert.NoError(t, err)
messageID := uint32(1)
err = binary.Write(mockConn.readBuf, binary.LittleEndian, messageID)
assert.NoError(t, err)
err = binary.Write(mockConn.readBuf, binary.LittleEndian, uint32(len(data)))
assert.NoError(t, err)
_, err = mockConn.readBuf.Write(data)
assert.NoError(t, err)
// Add EOF to signal end of stream
err = binary.Write(mockConn.readBuf, binary.LittleEndian, uint32(0))
assert.NoError(t, err)
err = binary.Write(mockConn.readBuf, binary.LittleEndian, uint32(0))
assert.NoError(t, err)
done := make(chan struct{})
go func() {
e.handleConnection(mockConn)
close(done)
}()
var receivedMessage *manager.ClientStreamMessage
select {
case receivedMessage = <-eventsChan:
case <-time.After(2 * time.Second):
t.Fatal("Timeout waiting for message in eventsChan")
}
assert.NotNil(t, receivedMessage)
select {
case <-done:
// handleConnection has exited
case <-time.After(2 * time.Second):
t.Fatal("Timeout waiting for handleConnection to exit")
}
// Check if ack was written
var receivedAck uint32
err = binary.Read(mockConn.writeBuf, binary.LittleEndian, &receivedAck)
assert.NoError(t, err)
assert.Equal(t, messageID, receivedAck)
// Ensure no unexpected calls were made on the mock
mockConn.AssertExpectations(t)
})
}
}
+249 -1518
View File
File diff suppressed because it is too large Load Diff
+18 -96
View File
@@ -3,40 +3,42 @@
syntax = "proto3";
import "google/protobuf/timestamp.proto";
import "google/protobuf/empty.proto";
package manager;
option go_package = "./manager";
service ManagerService {
rpc Process(stream ClientStreamMessage) returns (stream ServerStreamMessage) {}
rpc CreateVm(CreateReq) returns (CreateRes) {}
rpc RemoveVm(RemoveReq) returns (google.protobuf.Empty) {}
rpc SVMInfo(SVMInfoReq) returns (SVMInfoRes) {}
rpc AttestationPolicy(AttestationPolicyReq) returns (AttestationPolicyRes) {}
}
message Terminate {
string message = 1;
message CreateReq{
string agent_log_level = 1;
bytes agent_cvm_server_ca_cert = 2;
bytes agent_cvm_client_key = 3;
bytes agent_cvm_client_cert = 4;
string agent_cvm_server_url = 5;
}
message StopComputation {
string computation_id = 1;
message CreateRes{
string forwarded_port = 1;
string svm_id = 2;
}
message StopComputationResponse {
string computation_id = 1;
string message = 2;
message RemoveReq{
string svm_id = 1;
}
message RunResponse{
string agent_port = 1;
string computation_id = 2;
}
message AttestationPolicy{
message AttestationPolicyRes{
bytes info = 1;
string id = 2;
}
message SVMInfo{
message SVMInfoRes{
string id = 1;
string ovmf_version = 2;
int32 cpu_num = 3;
@@ -45,60 +47,6 @@ message SVMInfo{
string eos_version = 6;
}
message AgentEvent {
string event_type = 1;
google.protobuf.Timestamp timestamp = 2;
string computation_id = 3;
bytes details = 4;
string originator = 5;
string status = 6;
}
message AgentLog {
string message = 1;
string computation_id = 2;
string level = 3;
google.protobuf.Timestamp timestamp = 4;
}
message ClientStreamMessage {
oneof message {
AgentLog agent_log = 1;
AgentEvent agent_event = 2;
RunResponse run_res = 3;
AttestationPolicy attestationPolicy = 4;
StopComputationResponse stopComputationRes = 5;
SVMInfo svm_info = 6;
}
}
message ServerStreamMessage {
oneof message {
RunReqChunks runReqChunks = 1;
ComputationRunReq runReq = 2;
Terminate terminateReq = 3;
StopComputation stopComputation = 4;
AttestationPolicyReq attestationPolicyReq = 5;
SVMInfoReq svmInfoReq = 6;
}
}
message RunReqChunks {
bytes data = 1;
string id = 2;
bool is_last = 3;
}
message ComputationRunReq {
string id = 1;
string name = 2;
string description = 3;
repeated Dataset datasets = 4;
Algorithm algorithm = 5;
repeated ResultConsumer result_consumers = 6;
AgentConfig agent_config = 7;
}
message AttestationPolicyReq {
string id = 1;
}
@@ -107,29 +55,3 @@ message SVMInfoReq {
string id = 1;
}
message ResultConsumer {
bytes userKey = 1;
}
message Dataset {
bytes hash = 1; // should be sha3.Sum256, 32 byte length.
bytes userKey = 2;
string filename = 3;
}
message Algorithm {
bytes hash = 1; // should be sha3.Sum256, 32 byte length.
bytes userKey = 2;
}
message AgentConfig {
string port = 1;
string host = 2;
string cert_file = 3;
string key_file = 4;
string client_ca_file = 5;
string server_ca_file = 6;
string log_level = 7;
bool attested_tls = 8;
}
+143 -22
View File
@@ -4,7 +4,7 @@
// Code generated by protoc-gen-go-grpc. DO NOT EDIT.
// versions:
// - protoc-gen-go-grpc v1.5.1
// - protoc v5.28.1
// - protoc v5.29.0
// source: manager/manager.proto
package manager
@@ -14,6 +14,7 @@ import (
grpc "google.golang.org/grpc"
codes "google.golang.org/grpc/codes"
status "google.golang.org/grpc/status"
emptypb "google.golang.org/protobuf/types/known/emptypb"
)
// This is a compile-time assertion to ensure that this generated file
@@ -22,14 +23,20 @@ import (
const _ = grpc.SupportPackageIsVersion9
const (
ManagerService_Process_FullMethodName = "/manager.ManagerService/Process"
ManagerService_CreateVm_FullMethodName = "/manager.ManagerService/CreateVm"
ManagerService_RemoveVm_FullMethodName = "/manager.ManagerService/RemoveVm"
ManagerService_SVMInfo_FullMethodName = "/manager.ManagerService/SVMInfo"
ManagerService_AttestationPolicy_FullMethodName = "/manager.ManagerService/AttestationPolicy"
)
// ManagerServiceClient is the client API for ManagerService service.
//
// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream.
type ManagerServiceClient interface {
Process(ctx context.Context, opts ...grpc.CallOption) (grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage], error)
CreateVm(ctx context.Context, in *CreateReq, opts ...grpc.CallOption) (*CreateRes, error)
RemoveVm(ctx context.Context, in *RemoveReq, opts ...grpc.CallOption) (*emptypb.Empty, error)
SVMInfo(ctx context.Context, in *SVMInfoReq, opts ...grpc.CallOption) (*SVMInfoRes, error)
AttestationPolicy(ctx context.Context, in *AttestationPolicyReq, opts ...grpc.CallOption) (*AttestationPolicyRes, error)
}
type managerServiceClient struct {
@@ -40,24 +47,54 @@ func NewManagerServiceClient(cc grpc.ClientConnInterface) ManagerServiceClient {
return &managerServiceClient{cc}
}
func (c *managerServiceClient) Process(ctx context.Context, opts ...grpc.CallOption) (grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage], error) {
func (c *managerServiceClient) CreateVm(ctx context.Context, in *CreateReq, opts ...grpc.CallOption) (*CreateRes, error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
stream, err := c.cc.NewStream(ctx, &ManagerService_ServiceDesc.Streams[0], ManagerService_Process_FullMethodName, cOpts...)
out := new(CreateRes)
err := c.cc.Invoke(ctx, ManagerService_CreateVm_FullMethodName, in, out, cOpts...)
if err != nil {
return nil, err
}
x := &grpc.GenericClientStream[ClientStreamMessage, ServerStreamMessage]{ClientStream: stream}
return x, nil
return out, nil
}
// This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name.
type ManagerService_ProcessClient = grpc.BidiStreamingClient[ClientStreamMessage, ServerStreamMessage]
func (c *managerServiceClient) RemoveVm(ctx context.Context, in *RemoveReq, opts ...grpc.CallOption) (*emptypb.Empty, error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
out := new(emptypb.Empty)
err := c.cc.Invoke(ctx, ManagerService_RemoveVm_FullMethodName, in, out, cOpts...)
if err != nil {
return nil, err
}
return out, nil
}
func (c *managerServiceClient) SVMInfo(ctx context.Context, in *SVMInfoReq, opts ...grpc.CallOption) (*SVMInfoRes, error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
out := new(SVMInfoRes)
err := c.cc.Invoke(ctx, ManagerService_SVMInfo_FullMethodName, in, out, cOpts...)
if err != nil {
return nil, err
}
return out, nil
}
func (c *managerServiceClient) AttestationPolicy(ctx context.Context, in *AttestationPolicyReq, opts ...grpc.CallOption) (*AttestationPolicyRes, error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
out := new(AttestationPolicyRes)
err := c.cc.Invoke(ctx, ManagerService_AttestationPolicy_FullMethodName, in, out, cOpts...)
if err != nil {
return nil, err
}
return out, nil
}
// ManagerServiceServer is the server API for ManagerService service.
// All implementations must embed UnimplementedManagerServiceServer
// for forward compatibility.
type ManagerServiceServer interface {
Process(grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]) error
CreateVm(context.Context, *CreateReq) (*CreateRes, error)
RemoveVm(context.Context, *RemoveReq) (*emptypb.Empty, error)
SVMInfo(context.Context, *SVMInfoReq) (*SVMInfoRes, error)
AttestationPolicy(context.Context, *AttestationPolicyReq) (*AttestationPolicyRes, error)
mustEmbedUnimplementedManagerServiceServer()
}
@@ -68,8 +105,17 @@ type ManagerServiceServer interface {
// pointer dereference when methods are called.
type UnimplementedManagerServiceServer struct{}
func (UnimplementedManagerServiceServer) Process(grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]) error {
return status.Errorf(codes.Unimplemented, "method Process not implemented")
func (UnimplementedManagerServiceServer) CreateVm(context.Context, *CreateReq) (*CreateRes, error) {
return nil, status.Errorf(codes.Unimplemented, "method CreateVm not implemented")
}
func (UnimplementedManagerServiceServer) RemoveVm(context.Context, *RemoveReq) (*emptypb.Empty, error) {
return nil, status.Errorf(codes.Unimplemented, "method RemoveVm not implemented")
}
func (UnimplementedManagerServiceServer) SVMInfo(context.Context, *SVMInfoReq) (*SVMInfoRes, error) {
return nil, status.Errorf(codes.Unimplemented, "method SVMInfo not implemented")
}
func (UnimplementedManagerServiceServer) AttestationPolicy(context.Context, *AttestationPolicyReq) (*AttestationPolicyRes, error) {
return nil, status.Errorf(codes.Unimplemented, "method AttestationPolicy not implemented")
}
func (UnimplementedManagerServiceServer) mustEmbedUnimplementedManagerServiceServer() {}
func (UnimplementedManagerServiceServer) testEmbeddedByValue() {}
@@ -92,12 +138,77 @@ func RegisterManagerServiceServer(s grpc.ServiceRegistrar, srv ManagerServiceSer
s.RegisterService(&ManagerService_ServiceDesc, srv)
}
func _ManagerService_Process_Handler(srv interface{}, stream grpc.ServerStream) error {
return srv.(ManagerServiceServer).Process(&grpc.GenericServerStream[ClientStreamMessage, ServerStreamMessage]{ServerStream: stream})
func _ManagerService_CreateVm_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
in := new(CreateReq)
if err := dec(in); err != nil {
return nil, err
}
if interceptor == nil {
return srv.(ManagerServiceServer).CreateVm(ctx, in)
}
info := &grpc.UnaryServerInfo{
Server: srv,
FullMethod: ManagerService_CreateVm_FullMethodName,
}
handler := func(ctx context.Context, req interface{}) (interface{}, error) {
return srv.(ManagerServiceServer).CreateVm(ctx, req.(*CreateReq))
}
return interceptor(ctx, in, info, handler)
}
// This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name.
type ManagerService_ProcessServer = grpc.BidiStreamingServer[ClientStreamMessage, ServerStreamMessage]
func _ManagerService_RemoveVm_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
in := new(RemoveReq)
if err := dec(in); err != nil {
return nil, err
}
if interceptor == nil {
return srv.(ManagerServiceServer).RemoveVm(ctx, in)
}
info := &grpc.UnaryServerInfo{
Server: srv,
FullMethod: ManagerService_RemoveVm_FullMethodName,
}
handler := func(ctx context.Context, req interface{}) (interface{}, error) {
return srv.(ManagerServiceServer).RemoveVm(ctx, req.(*RemoveReq))
}
return interceptor(ctx, in, info, handler)
}
func _ManagerService_SVMInfo_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
in := new(SVMInfoReq)
if err := dec(in); err != nil {
return nil, err
}
if interceptor == nil {
return srv.(ManagerServiceServer).SVMInfo(ctx, in)
}
info := &grpc.UnaryServerInfo{
Server: srv,
FullMethod: ManagerService_SVMInfo_FullMethodName,
}
handler := func(ctx context.Context, req interface{}) (interface{}, error) {
return srv.(ManagerServiceServer).SVMInfo(ctx, req.(*SVMInfoReq))
}
return interceptor(ctx, in, info, handler)
}
func _ManagerService_AttestationPolicy_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
in := new(AttestationPolicyReq)
if err := dec(in); err != nil {
return nil, err
}
if interceptor == nil {
return srv.(ManagerServiceServer).AttestationPolicy(ctx, in)
}
info := &grpc.UnaryServerInfo{
Server: srv,
FullMethod: ManagerService_AttestationPolicy_FullMethodName,
}
handler := func(ctx context.Context, req interface{}) (interface{}, error) {
return srv.(ManagerServiceServer).AttestationPolicy(ctx, req.(*AttestationPolicyReq))
}
return interceptor(ctx, in, info, handler)
}
// ManagerService_ServiceDesc is the grpc.ServiceDesc for ManagerService service.
// It's only intended for direct use with grpc.RegisterService,
@@ -105,14 +216,24 @@ type ManagerService_ProcessServer = grpc.BidiStreamingServer[ClientStreamMessage
var ManagerService_ServiceDesc = grpc.ServiceDesc{
ServiceName: "manager.ManagerService",
HandlerType: (*ManagerServiceServer)(nil),
Methods: []grpc.MethodDesc{},
Streams: []grpc.StreamDesc{
Methods: []grpc.MethodDesc{
{
StreamName: "Process",
Handler: _ManagerService_Process_Handler,
ServerStreams: true,
ClientStreams: true,
MethodName: "CreateVm",
Handler: _ManagerService_CreateVm_Handler,
},
{
MethodName: "RemoveVm",
Handler: _ManagerService_RemoveVm_Handler,
},
{
MethodName: "SVMInfo",
Handler: _ManagerService_SVMInfo_Handler,
},
{
MethodName: "AttestationPolicy",
Handler: _ManagerService_AttestationPolicy_Handler,
},
},
Streams: []grpc.StreamDesc{},
Metadata: "manager/manager.proto",
}
-68
View File
@@ -1,68 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package manager_test
import (
"bytes"
"context"
"testing"
"github.com/ultravioletrs/cocos/manager"
"google.golang.org/grpc"
"google.golang.org/protobuf/proto"
)
func TestProcess(t *testing.T) {
ctx := context.Background()
conn, err := grpc.DialContext(ctx, "bufnet", grpc.WithContextDialer(bufDialer), grpc.WithInsecure())
if err != nil {
t.Fatalf("Failed to dial bufnet: %v", err)
}
defer conn.Close()
client := manager.NewManagerServiceClient(conn)
stream, err := client.Process(ctx)
if err != nil {
t.Fatalf("Process failed: %v", err)
}
var data bytes.Buffer
for {
msg, err := stream.Recv()
if err != nil {
t.Fatalf("Failed to receive ServerStreamMessage: %v", err)
}
switch m := msg.Message.(type) {
case *manager.ServerStreamMessage_TerminateReq:
if m.TerminateReq.Message != "test terminate" {
t.Fatalf("Unexpected terminate message: %v", m.TerminateReq.Message)
}
case *manager.ServerStreamMessage_RunReqChunks:
if len(m.RunReqChunks.Data) == 0 {
var runReq manager.ComputationRunReq
if err = proto.Unmarshal(data.Bytes(), &runReq); err != nil {
t.Fatalf("Failed to create run request: %v", err)
}
runRes := &manager.ClientStreamMessage_AgentLog{
AgentLog: &manager.AgentLog{
Message: "test log",
ComputationId: "comp1",
Level: "DEBUG",
},
}
if runReq.Id != "1" || runReq.Name != "sample computation" || runReq.Description != "sample description" {
t.Fatalf("Unexpected run request message: %v", &runReq)
}
if err := stream.Send(&manager.ClientStreamMessage{Message: runRes}); err != nil {
t.Fatalf("Failed to send ClientStreamMessage: %v", err)
}
return
}
data.Write(m.RunReqChunks.Data)
default:
t.Fatalf("Unexpected message type: %T", m)
}
}
}
+340
View File
@@ -0,0 +1,340 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
// Code generated by mockery v2.43.2. DO NOT EDIT.
package mocks
import (
context "context"
grpc "google.golang.org/grpc"
emptypb "google.golang.org/protobuf/types/known/emptypb"
manager "github.com/ultravioletrs/cocos/manager"
mock "github.com/stretchr/testify/mock"
)
// ManagerServiceClient is an autogenerated mock type for the ManagerServiceClient type
type ManagerServiceClient struct {
mock.Mock
}
type ManagerServiceClient_Expecter struct {
mock *mock.Mock
}
func (_m *ManagerServiceClient) EXPECT() *ManagerServiceClient_Expecter {
return &ManagerServiceClient_Expecter{mock: &_m.Mock}
}
// AttestationPolicy provides a mock function with given fields: ctx, in, opts
func (_m *ManagerServiceClient) AttestationPolicy(ctx context.Context, in *manager.AttestationPolicyReq, opts ...grpc.CallOption) (*manager.AttestationPolicyRes, error) {
_va := make([]interface{}, len(opts))
for _i := range opts {
_va[_i] = opts[_i]
}
var _ca []interface{}
_ca = append(_ca, ctx, in)
_ca = append(_ca, _va...)
ret := _m.Called(_ca...)
if len(ret) == 0 {
panic("no return value specified for AttestationPolicy")
}
var r0 *manager.AttestationPolicyRes
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.AttestationPolicyReq, ...grpc.CallOption) (*manager.AttestationPolicyRes, error)); ok {
return rf(ctx, in, opts...)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.AttestationPolicyReq, ...grpc.CallOption) *manager.AttestationPolicyRes); ok {
r0 = rf(ctx, in, opts...)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*manager.AttestationPolicyRes)
}
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.AttestationPolicyReq, ...grpc.CallOption) error); ok {
r1 = rf(ctx, in, opts...)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// ManagerServiceClient_AttestationPolicy_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AttestationPolicy'
type ManagerServiceClient_AttestationPolicy_Call struct {
*mock.Call
}
// AttestationPolicy is a helper method to define mock.On call
// - ctx context.Context
// - in *manager.AttestationPolicyReq
// - opts ...grpc.CallOption
func (_e *ManagerServiceClient_Expecter) AttestationPolicy(ctx interface{}, in interface{}, opts ...interface{}) *ManagerServiceClient_AttestationPolicy_Call {
return &ManagerServiceClient_AttestationPolicy_Call{Call: _e.mock.On("AttestationPolicy",
append([]interface{}{ctx, in}, opts...)...)}
}
func (_c *ManagerServiceClient_AttestationPolicy_Call) Run(run func(ctx context.Context, in *manager.AttestationPolicyReq, opts ...grpc.CallOption)) *ManagerServiceClient_AttestationPolicy_Call {
_c.Call.Run(func(args mock.Arguments) {
variadicArgs := make([]grpc.CallOption, len(args)-2)
for i, a := range args[2:] {
if a != nil {
variadicArgs[i] = a.(grpc.CallOption)
}
}
run(args[0].(context.Context), args[1].(*manager.AttestationPolicyReq), variadicArgs...)
})
return _c
}
func (_c *ManagerServiceClient_AttestationPolicy_Call) Return(_a0 *manager.AttestationPolicyRes, _a1 error) *ManagerServiceClient_AttestationPolicy_Call {
_c.Call.Return(_a0, _a1)
return _c
}
func (_c *ManagerServiceClient_AttestationPolicy_Call) RunAndReturn(run func(context.Context, *manager.AttestationPolicyReq, ...grpc.CallOption) (*manager.AttestationPolicyRes, error)) *ManagerServiceClient_AttestationPolicy_Call {
_c.Call.Return(run)
return _c
}
// CreateVm provides a mock function with given fields: ctx, in, opts
func (_m *ManagerServiceClient) CreateVm(ctx context.Context, in *manager.CreateReq, opts ...grpc.CallOption) (*manager.CreateRes, error) {
_va := make([]interface{}, len(opts))
for _i := range opts {
_va[_i] = opts[_i]
}
var _ca []interface{}
_ca = append(_ca, ctx, in)
_ca = append(_ca, _va...)
ret := _m.Called(_ca...)
if len(ret) == 0 {
panic("no return value specified for CreateVm")
}
var r0 *manager.CreateRes
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.CreateReq, ...grpc.CallOption) (*manager.CreateRes, error)); ok {
return rf(ctx, in, opts...)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.CreateReq, ...grpc.CallOption) *manager.CreateRes); ok {
r0 = rf(ctx, in, opts...)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*manager.CreateRes)
}
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.CreateReq, ...grpc.CallOption) error); ok {
r1 = rf(ctx, in, opts...)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// ManagerServiceClient_CreateVm_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateVm'
type ManagerServiceClient_CreateVm_Call struct {
*mock.Call
}
// CreateVm is a helper method to define mock.On call
// - ctx context.Context
// - in *manager.CreateReq
// - opts ...grpc.CallOption
func (_e *ManagerServiceClient_Expecter) CreateVm(ctx interface{}, in interface{}, opts ...interface{}) *ManagerServiceClient_CreateVm_Call {
return &ManagerServiceClient_CreateVm_Call{Call: _e.mock.On("CreateVm",
append([]interface{}{ctx, in}, opts...)...)}
}
func (_c *ManagerServiceClient_CreateVm_Call) Run(run func(ctx context.Context, in *manager.CreateReq, opts ...grpc.CallOption)) *ManagerServiceClient_CreateVm_Call {
_c.Call.Run(func(args mock.Arguments) {
variadicArgs := make([]grpc.CallOption, len(args)-2)
for i, a := range args[2:] {
if a != nil {
variadicArgs[i] = a.(grpc.CallOption)
}
}
run(args[0].(context.Context), args[1].(*manager.CreateReq), variadicArgs...)
})
return _c
}
func (_c *ManagerServiceClient_CreateVm_Call) Return(_a0 *manager.CreateRes, _a1 error) *ManagerServiceClient_CreateVm_Call {
_c.Call.Return(_a0, _a1)
return _c
}
func (_c *ManagerServiceClient_CreateVm_Call) RunAndReturn(run func(context.Context, *manager.CreateReq, ...grpc.CallOption) (*manager.CreateRes, error)) *ManagerServiceClient_CreateVm_Call {
_c.Call.Return(run)
return _c
}
// RemoveVm provides a mock function with given fields: ctx, in, opts
func (_m *ManagerServiceClient) RemoveVm(ctx context.Context, in *manager.RemoveReq, opts ...grpc.CallOption) (*emptypb.Empty, error) {
_va := make([]interface{}, len(opts))
for _i := range opts {
_va[_i] = opts[_i]
}
var _ca []interface{}
_ca = append(_ca, ctx, in)
_ca = append(_ca, _va...)
ret := _m.Called(_ca...)
if len(ret) == 0 {
panic("no return value specified for RemoveVm")
}
var r0 *emptypb.Empty
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.RemoveReq, ...grpc.CallOption) (*emptypb.Empty, error)); ok {
return rf(ctx, in, opts...)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.RemoveReq, ...grpc.CallOption) *emptypb.Empty); ok {
r0 = rf(ctx, in, opts...)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*emptypb.Empty)
}
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.RemoveReq, ...grpc.CallOption) error); ok {
r1 = rf(ctx, in, opts...)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// ManagerServiceClient_RemoveVm_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveVm'
type ManagerServiceClient_RemoveVm_Call struct {
*mock.Call
}
// RemoveVm is a helper method to define mock.On call
// - ctx context.Context
// - in *manager.RemoveReq
// - opts ...grpc.CallOption
func (_e *ManagerServiceClient_Expecter) RemoveVm(ctx interface{}, in interface{}, opts ...interface{}) *ManagerServiceClient_RemoveVm_Call {
return &ManagerServiceClient_RemoveVm_Call{Call: _e.mock.On("RemoveVm",
append([]interface{}{ctx, in}, opts...)...)}
}
func (_c *ManagerServiceClient_RemoveVm_Call) Run(run func(ctx context.Context, in *manager.RemoveReq, opts ...grpc.CallOption)) *ManagerServiceClient_RemoveVm_Call {
_c.Call.Run(func(args mock.Arguments) {
variadicArgs := make([]grpc.CallOption, len(args)-2)
for i, a := range args[2:] {
if a != nil {
variadicArgs[i] = a.(grpc.CallOption)
}
}
run(args[0].(context.Context), args[1].(*manager.RemoveReq), variadicArgs...)
})
return _c
}
func (_c *ManagerServiceClient_RemoveVm_Call) Return(_a0 *emptypb.Empty, _a1 error) *ManagerServiceClient_RemoveVm_Call {
_c.Call.Return(_a0, _a1)
return _c
}
func (_c *ManagerServiceClient_RemoveVm_Call) RunAndReturn(run func(context.Context, *manager.RemoveReq, ...grpc.CallOption) (*emptypb.Empty, error)) *ManagerServiceClient_RemoveVm_Call {
_c.Call.Return(run)
return _c
}
// SVMInfo provides a mock function with given fields: ctx, in, opts
func (_m *ManagerServiceClient) SVMInfo(ctx context.Context, in *manager.SVMInfoReq, opts ...grpc.CallOption) (*manager.SVMInfoRes, error) {
_va := make([]interface{}, len(opts))
for _i := range opts {
_va[_i] = opts[_i]
}
var _ca []interface{}
_ca = append(_ca, ctx, in)
_ca = append(_ca, _va...)
ret := _m.Called(_ca...)
if len(ret) == 0 {
panic("no return value specified for SVMInfo")
}
var r0 *manager.SVMInfoRes
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.SVMInfoReq, ...grpc.CallOption) (*manager.SVMInfoRes, error)); ok {
return rf(ctx, in, opts...)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.SVMInfoReq, ...grpc.CallOption) *manager.SVMInfoRes); ok {
r0 = rf(ctx, in, opts...)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*manager.SVMInfoRes)
}
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.SVMInfoReq, ...grpc.CallOption) error); ok {
r1 = rf(ctx, in, opts...)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// ManagerServiceClient_SVMInfo_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SVMInfo'
type ManagerServiceClient_SVMInfo_Call struct {
*mock.Call
}
// SVMInfo is a helper method to define mock.On call
// - ctx context.Context
// - in *manager.SVMInfoReq
// - opts ...grpc.CallOption
func (_e *ManagerServiceClient_Expecter) SVMInfo(ctx interface{}, in interface{}, opts ...interface{}) *ManagerServiceClient_SVMInfo_Call {
return &ManagerServiceClient_SVMInfo_Call{Call: _e.mock.On("SVMInfo",
append([]interface{}{ctx, in}, opts...)...)}
}
func (_c *ManagerServiceClient_SVMInfo_Call) Run(run func(ctx context.Context, in *manager.SVMInfoReq, opts ...grpc.CallOption)) *ManagerServiceClient_SVMInfo_Call {
_c.Call.Run(func(args mock.Arguments) {
variadicArgs := make([]grpc.CallOption, len(args)-2)
for i, a := range args[2:] {
if a != nil {
variadicArgs[i] = a.(grpc.CallOption)
}
}
run(args[0].(context.Context), args[1].(*manager.SVMInfoReq), variadicArgs...)
})
return _c
}
func (_c *ManagerServiceClient_SVMInfo_Call) Return(_a0 *manager.SVMInfoRes, _a1 error) *ManagerServiceClient_SVMInfo_Call {
_c.Call.Return(_a0, _a1)
return _c
}
func (_c *ManagerServiceClient_SVMInfo_Call) RunAndReturn(run func(context.Context, *manager.SVMInfoReq, ...grpc.CallOption) (*manager.SVMInfoRes, error)) *ManagerServiceClient_SVMInfo_Call {
_c.Call.Return(run)
return _c
}
// NewManagerServiceClient creates a new instance of ManagerServiceClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations.
// The first argument is typically a *testing.T value.
func NewManagerServiceClient(t interface {
mock.TestingT
Cleanup(func())
}) *ManagerServiceClient {
mock := &ManagerServiceClient{}
mock.Mock.Test(t)
t.Cleanup(func() { mock.AssertExpectations(t) })
return mock
}
+92 -118
View File
@@ -25,6 +25,70 @@ func (_m *Service) EXPECT() *Service_Expecter {
return &Service_Expecter{mock: &_m.Mock}
}
// CreateVM provides a mock function with given fields: ctx, req
func (_m *Service) CreateVM(ctx context.Context, req *manager.CreateReq) (string, string, error) {
ret := _m.Called(ctx, req)
if len(ret) == 0 {
panic("no return value specified for CreateVM")
}
var r0 string
var r1 string
var r2 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.CreateReq) (string, string, error)); ok {
return rf(ctx, req)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.CreateReq) string); ok {
r0 = rf(ctx, req)
} else {
r0 = ret.Get(0).(string)
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.CreateReq) string); ok {
r1 = rf(ctx, req)
} else {
r1 = ret.Get(1).(string)
}
if rf, ok := ret.Get(2).(func(context.Context, *manager.CreateReq) error); ok {
r2 = rf(ctx, req)
} else {
r2 = ret.Error(2)
}
return r0, r1, r2
}
// Service_CreateVM_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateVM'
type Service_CreateVM_Call struct {
*mock.Call
}
// CreateVM is a helper method to define mock.On call
// - ctx context.Context
// - req *manager.CreateReq
func (_e *Service_Expecter) CreateVM(ctx interface{}, req interface{}) *Service_CreateVM_Call {
return &Service_CreateVM_Call{Call: _e.mock.On("CreateVM", ctx, req)}
}
func (_c *Service_CreateVM_Call) Run(run func(ctx context.Context, req *manager.CreateReq)) *Service_CreateVM_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context), args[1].(*manager.CreateReq))
})
return _c
}
func (_c *Service_CreateVM_Call) Return(_a0 string, _a1 string, _a2 error) *Service_CreateVM_Call {
_c.Call.Return(_a0, _a1, _a2)
return _c
}
func (_c *Service_CreateVM_Call) RunAndReturn(run func(context.Context, *manager.CreateReq) (string, string, error)) *Service_CreateVM_Call {
_c.Call.Return(run)
return _c
}
// FetchAttestationPolicy provides a mock function with given fields: ctx, computationID
func (_m *Service) FetchAttestationPolicy(ctx context.Context, computationID string) ([]byte, error) {
ret := _m.Called(ctx, computationID)
@@ -84,35 +148,49 @@ func (_c *Service_FetchAttestationPolicy_Call) RunAndReturn(run func(context.Con
return _c
}
// ReportBrokenConnection provides a mock function with given fields: addr
func (_m *Service) ReportBrokenConnection(addr string) {
_m.Called(addr)
// RemoveVM provides a mock function with given fields: ctx, computationID
func (_m *Service) RemoveVM(ctx context.Context, computationID string) error {
ret := _m.Called(ctx, computationID)
if len(ret) == 0 {
panic("no return value specified for RemoveVM")
}
var r0 error
if rf, ok := ret.Get(0).(func(context.Context, string) error); ok {
r0 = rf(ctx, computationID)
} else {
r0 = ret.Error(0)
}
return r0
}
// Service_ReportBrokenConnection_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReportBrokenConnection'
type Service_ReportBrokenConnection_Call struct {
// Service_RemoveVM_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveVM'
type Service_RemoveVM_Call struct {
*mock.Call
}
// ReportBrokenConnection is a helper method to define mock.On call
// - addr string
func (_e *Service_Expecter) ReportBrokenConnection(addr interface{}) *Service_ReportBrokenConnection_Call {
return &Service_ReportBrokenConnection_Call{Call: _e.mock.On("ReportBrokenConnection", addr)}
// RemoveVM is a helper method to define mock.On call
// - ctx context.Context
// - computationID string
func (_e *Service_Expecter) RemoveVM(ctx interface{}, computationID interface{}) *Service_RemoveVM_Call {
return &Service_RemoveVM_Call{Call: _e.mock.On("RemoveVM", ctx, computationID)}
}
func (_c *Service_ReportBrokenConnection_Call) Run(run func(addr string)) *Service_ReportBrokenConnection_Call {
func (_c *Service_RemoveVM_Call) Run(run func(ctx context.Context, computationID string)) *Service_RemoveVM_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(string))
run(args[0].(context.Context), args[1].(string))
})
return _c
}
func (_c *Service_ReportBrokenConnection_Call) Return() *Service_ReportBrokenConnection_Call {
_c.Call.Return()
func (_c *Service_RemoveVM_Call) Return(_a0 error) *Service_RemoveVM_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *Service_ReportBrokenConnection_Call) RunAndReturn(run func(string)) *Service_ReportBrokenConnection_Call {
func (_c *Service_RemoveVM_Call) RunAndReturn(run func(context.Context, string) error) *Service_RemoveVM_Call {
_c.Call.Return(run)
return _c
}
@@ -187,110 +265,6 @@ func (_c *Service_ReturnSVMInfo_Call) RunAndReturn(run func(context.Context) (st
return _c
}
// Run provides a mock function with given fields: ctx, c
func (_m *Service) Run(ctx context.Context, c *manager.ComputationRunReq) (string, error) {
ret := _m.Called(ctx, c)
if len(ret) == 0 {
panic("no return value specified for Run")
}
var r0 string
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, *manager.ComputationRunReq) (string, error)); ok {
return rf(ctx, c)
}
if rf, ok := ret.Get(0).(func(context.Context, *manager.ComputationRunReq) string); ok {
r0 = rf(ctx, c)
} else {
r0 = ret.Get(0).(string)
}
if rf, ok := ret.Get(1).(func(context.Context, *manager.ComputationRunReq) error); ok {
r1 = rf(ctx, c)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// Service_Run_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Run'
type Service_Run_Call struct {
*mock.Call
}
// Run is a helper method to define mock.On call
// - ctx context.Context
// - c *manager.ComputationRunReq
func (_e *Service_Expecter) Run(ctx interface{}, c interface{}) *Service_Run_Call {
return &Service_Run_Call{Call: _e.mock.On("Run", ctx, c)}
}
func (_c *Service_Run_Call) Run(run func(ctx context.Context, c *manager.ComputationRunReq)) *Service_Run_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context), args[1].(*manager.ComputationRunReq))
})
return _c
}
func (_c *Service_Run_Call) Return(_a0 string, _a1 error) *Service_Run_Call {
_c.Call.Return(_a0, _a1)
return _c
}
func (_c *Service_Run_Call) RunAndReturn(run func(context.Context, *manager.ComputationRunReq) (string, error)) *Service_Run_Call {
_c.Call.Return(run)
return _c
}
// Stop provides a mock function with given fields: ctx, computationID
func (_m *Service) Stop(ctx context.Context, computationID string) error {
ret := _m.Called(ctx, computationID)
if len(ret) == 0 {
panic("no return value specified for Stop")
}
var r0 error
if rf, ok := ret.Get(0).(func(context.Context, string) error); ok {
r0 = rf(ctx, computationID)
} else {
r0 = ret.Error(0)
}
return r0
}
// Service_Stop_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Stop'
type Service_Stop_Call struct {
*mock.Call
}
// Stop is a helper method to define mock.On call
// - ctx context.Context
// - computationID string
func (_e *Service_Expecter) Stop(ctx interface{}, computationID interface{}) *Service_Stop_Call {
return &Service_Stop_Call{Call: _e.mock.On("Stop", ctx, computationID)}
}
func (_c *Service_Stop_Call) Run(run func(ctx context.Context, computationID string)) *Service_Stop_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(context.Context), args[1].(string))
})
return _c
}
func (_c *Service_Stop_Call) Return(_a0 error) *Service_Stop_Call {
_c.Call.Return(_a0)
return _c
}
func (_c *Service_Stop_Call) RunAndReturn(run func(context.Context, string) error) *Service_Stop_Call {
_c.Call.Return(run)
return _c
}
// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations.
// The first argument is typically a *testing.T value.
func NewService(t interface {
+14
View File
@@ -106,6 +106,10 @@ type Config struct {
// ports
HostFwdRange string `env:"HOST_FWD_RANGE" envDefault:"6100-6200"`
// mounts
CertsMount string `env:"CERTS_MOUNT" envDefault:""`
EnvMount string `env:"ENV_MOUNT" envDefault:""`
}
func (config Config) ConstructQemuArgs() []string {
@@ -216,5 +220,15 @@ func (config Config) ConstructQemuArgs() []string {
args = append(args, "-monitor", config.Monitor)
if config.CertsMount != "" {
args = append(args, "-fsdev", fmt.Sprintf("local,id=cert_fs,path=%s,security_model=mapped", config.CertsMount))
args = append(args, "-device", "virtio-9p-pci,fsdev=cert_fs,mount_tag=certs_share")
}
if config.EnvMount != "" {
args = append(args, "-fsdev", fmt.Sprintf("local,id=env_fs,path=%s,security_model=mapped", config.EnvMount))
args = append(args, "-device", "virtio-9p-pci,fsdev=env_fs,mount_tag=env_share")
}
return args
}
+1 -1
View File
@@ -13,7 +13,7 @@ const jsonExt = ".json"
type VMState struct {
ID string
Config Config
VMinfo VMInfo
PID int
}
+5 -5
View File
@@ -29,7 +29,7 @@ func TestSaveVM(t *testing.T) {
state := VMState{
ID: "test-vm",
Config: Config{},
VMinfo: VMInfo{Config: Config{}},
PID: 1234,
}
@@ -50,8 +50,8 @@ func TestLoadVMs(t *testing.T) {
// Save two VMs
states := []VMState{
{ID: "vm1", Config: Config{}, PID: 1234},
{ID: "vm2", Config: Config{}, PID: 5678},
{ID: "vm1", VMinfo: VMInfo{Config: Config{}}, PID: 1234},
{ID: "vm2", VMinfo: VMInfo{Config: Config{}}, PID: 5678},
}
for _, state := range states {
@@ -82,7 +82,7 @@ func TestDeleteVM(t *testing.T) {
tempDir := t.TempDir()
fp, _ := NewFilePersistence(tempDir)
state := VMState{ID: "test-vm", Config: Config{}, PID: 1234}
state := VMState{ID: "test-vm", VMinfo: VMInfo{Config: Config{}}, PID: 1234}
// Save VM
if err := fp.SaveVM(state); err != nil {
@@ -126,7 +126,7 @@ func TestConcurrentAccess(t *testing.T) {
for i := 0; i < numGoroutines; i++ {
go func(id int) {
defer wg.Done()
state := VMState{ID: fmt.Sprintf("vm-%d", id), Config: Config{}, PID: id}
state := VMState{ID: fmt.Sprintf("vm-%d", id), VMinfo: VMInfo{Config: Config{}}, PID: id}
if err := fp.SaveVM(state); err != nil {
t.Errorf("Concurrent SaveVM failed: %v", err)
}
+37 -40
View File
@@ -13,7 +13,6 @@ import (
"github.com/ultravioletrs/cocos/internal"
"github.com/ultravioletrs/cocos/manager/vm"
"github.com/ultravioletrs/cocos/pkg/manager"
"google.golang.org/protobuf/types/known/timestamppb"
)
const (
@@ -25,20 +24,23 @@ const (
shutdownTimeout = 30 * time.Second
)
type VMInfo struct {
Config Config
LaunchTCB uint64 `env:"LAUNCH_TCB" envDefault:"0"`
}
type qemuVM struct {
config Config
cmd *exec.Cmd
eventsLogsSender vm.EventSender
computationId string
vmi VMInfo
cmd *exec.Cmd
computationId string
vm.StateMachine
}
func NewVM(config interface{}, eventsLogsSender vm.EventSender, computationId string) vm.VM {
func NewVM(config interface{}, computationId string) vm.VM {
return &qemuVM{
config: config.(Config),
eventsLogsSender: eventsLogsSender,
computationId: computationId,
StateMachine: vm.NewStateMachine(),
vmi: config.(VMInfo),
computationId: computationId,
StateMachine: vm.NewStateMachine(),
}
}
@@ -54,18 +56,18 @@ func (v *qemuVM) Start() (err error) {
return err
}
v.config.NetDevConfig.ID = fmt.Sprintf("%s-%s", v.config.NetDevConfig.ID, id)
v.config.SevConfig.ID = fmt.Sprintf("%s-%s", v.config.SevConfig.ID, id)
v.vmi.Config.NetDevConfig.ID = fmt.Sprintf("%s-%s", v.vmi.Config.NetDevConfig.ID, id)
v.vmi.Config.SevConfig.ID = fmt.Sprintf("%s-%s", v.vmi.Config.SevConfig.ID, id)
if !v.config.KernelHash {
if !v.vmi.Config.KernelHash {
// Copy firmware vars file.
srcFile := v.config.OVMFVarsConfig.File
srcFile := v.vmi.Config.OVMFVarsConfig.File
dstFile := fmt.Sprintf("%s/%s-%s.fd", tmpDir, firmwareVars, id)
err = internal.CopyFile(srcFile, dstFile)
if err != nil {
return err
}
v.config.OVMFVarsConfig.File = dstFile
v.vmi.Config.OVMFVarsConfig.File = dstFile
}
exe, args, err := v.executableAndArgs()
@@ -74,8 +76,8 @@ func (v *qemuVM) Start() (err error) {
}
v.cmd = exec.Command(exe, args...)
v.cmd.Stdout = &vm.Stdout{ComputationId: v.computationId, EventSender: v.eventsLogsSender}
v.cmd.Stderr = &vm.Stderr{EventSender: v.eventsLogsSender, ComputationId: v.computationId, StateMachine: v.StateMachine}
v.cmd.Stdout = os.Stdout
v.cmd.Stderr = os.Stderr
return v.cmd.Start()
}
@@ -84,15 +86,7 @@ func (v *qemuVM) Stop() error {
defer func() {
err := v.StateMachine.Transition(manager.StopComputationRun)
if err != nil {
if err := v.eventsLogsSender(&vm.Event{
EventType: v.StateMachine.State(),
Timestamp: timestamppb.Now(),
ComputationId: v.computationId,
Originator: "manager",
Status: manager.Warning.String(),
}); err != nil {
return
}
return
}
}()
err := v.cmd.Process.Signal(syscall.SIGTERM)
@@ -100,6 +94,18 @@ func (v *qemuVM) Stop() error {
return fmt.Errorf("failed to send SIGTERM: %v", err)
}
if v.vmi.Config.CertsMount != "" {
if err := os.RemoveAll(v.vmi.Config.CertsMount); err != nil {
return fmt.Errorf("failed to remove certs mount: %v", err)
}
}
if v.vmi.Config.EnvMount != "" {
if err := os.RemoveAll(v.vmi.Config.EnvMount); err != nil {
return fmt.Errorf("failed to remove env mount: %v", err)
}
}
done := make(chan error, 1)
go func() {
_, err := v.cmd.Process.Wait()
@@ -140,14 +146,14 @@ func (v *qemuVM) GetProcess() int {
}
func (v *qemuVM) executableAndArgs() (string, []string, error) {
exe, err := exec.LookPath(v.config.QemuBinPath)
exe, err := exec.LookPath(v.vmi.Config.QemuBinPath)
if err != nil {
return "", nil, err
}
args := v.config.ConstructQemuArgs()
args := v.vmi.Config.ConstructQemuArgs()
if v.config.UseSudo {
if v.vmi.Config.UseSudo {
args = append([]string{exe}, args...)
exe = "sudo"
}
@@ -158,15 +164,6 @@ func (v *qemuVM) executableAndArgs() (string, []string, error) {
func (v *qemuVM) checkVMProcessPeriodically() {
for {
if !processExists(v.GetProcess()) {
if err := v.eventsLogsSender(&vm.Event{
EventType: v.StateMachine.State(),
Timestamp: timestamppb.Now(),
ComputationId: v.computationId,
Originator: "manager",
Status: manager.Stopped.String(),
}); err != nil {
return
}
break
}
time.Sleep(interval)
@@ -191,9 +188,9 @@ func processExists(pid int) bool {
}
func (v *qemuVM) GetCID() int {
return v.config.GuestCID
return v.vmi.Config.GuestCID
}
func (v *qemuVM) GetConfig() interface{} {
return v.config
return v.vmi
}
+20 -49
View File
@@ -6,10 +6,8 @@ import (
"os"
"os/exec"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/ultravioletrs/cocos/manager/vm"
"github.com/ultravioletrs/cocos/manager/vm/mocks"
pkgmanager "github.com/ultravioletrs/cocos/pkg/manager"
)
@@ -17,9 +15,9 @@ import (
const testComputationID = "test-computation"
func TestNewVM(t *testing.T) {
config := Config{}
config := VMInfo{Config: Config{}}
vm := NewVM(config, func(event interface{}) error { return nil }, testComputationID)
vm := NewVM(config, testComputationID)
assert.NotNil(t, vm)
assert.IsType(t, &qemuVM{}, vm)
@@ -31,14 +29,14 @@ func TestStart(t *testing.T) {
assert.NoError(t, err)
defer os.Remove(tmpFile.Name())
config := Config{
config := VMInfo{Config: Config{
OVMFVarsConfig: OVMFVarsConfig{
File: tmpFile.Name(),
},
QemuBinPath: "echo",
}
}}
vm := NewVM(config, func(event interface{}) error { return nil }, testComputationID).(*qemuVM)
vm := NewVM(config, testComputationID).(*qemuVM)
err = vm.Start()
assert.NoError(t, err)
@@ -53,15 +51,15 @@ func TestStartSudo(t *testing.T) {
assert.NoError(t, err)
defer os.Remove(tmpFile.Name())
config := Config{
config := VMInfo{Config: Config{
OVMFVarsConfig: OVMFVarsConfig{
File: tmpFile.Name(),
},
QemuBinPath: "echo",
UseSudo: true,
}
}}
vm := NewVM(config, func(event interface{}) error { return nil }, testComputationID).(*qemuVM)
vm := NewVM(config, testComputationID).(*qemuVM)
err = vm.Start()
assert.NoError(t, err)
@@ -101,9 +99,6 @@ func TestStop(t *testing.T) {
Process: cmd.Process,
},
StateMachine: sm,
eventsLogsSender: func(event interface{}) error {
return nil
},
}
err = vm.Stop()
@@ -113,8 +108,8 @@ func TestStop(t *testing.T) {
func TestSetProcess(t *testing.T) {
vm := &qemuVM{
config: Config{
QemuBinPath: "echo", // Use 'echo' as a dummy QEMU binary
vmi: VMInfo{
Config: Config{QemuBinPath: "echo"}, // Use 'echo' as a dummy QEMU binary
},
}
@@ -139,9 +134,11 @@ func TestGetProcess(t *testing.T) {
func TestGetCID(t *testing.T) {
expectedCID := 42
vm := &qemuVM{
config: Config{
VSockConfig: VSockConfig{
GuestCID: expectedCID,
vmi: VMInfo{
Config: Config{
VSockConfig: VSockConfig{
GuestCID: expectedCID,
},
},
},
}
@@ -151,41 +148,15 @@ func TestGetCID(t *testing.T) {
}
func TestGetConfig(t *testing.T) {
expectedConfig := Config{
QemuBinPath: "echo",
expectedConfig := VMInfo{
Config: Config{
QemuBinPath: "echo",
},
}
vm := &qemuVM{
config: expectedConfig,
vmi: expectedConfig,
}
config := vm.GetConfig()
assert.Equal(t, expectedConfig, config)
}
func TestCheckVMProcessPeriodically(t *testing.T) {
logsChan := make(chan interface{}, 1)
vmi := &qemuVM{
eventsLogsSender: func(event interface{}) error {
logsChan <- event
return nil
},
computationId: testComputationID,
cmd: &exec.Cmd{
Process: &os.Process{Pid: -1}, // Use an invalid PID to simulate a stopped process
},
StateMachine: vm.NewStateMachine(),
}
go vmi.checkVMProcessPeriodically()
select {
case msg := <-logsChan:
assert.NotNil(t, msg)
msgE := msg.(*vm.Event)
assert.Equal(t, testComputationID, msgE.ComputationId)
assert.Equal(t, pkgmanager.VmProvision.String(), msgE.EventType)
assert.Equal(t, pkgmanager.Stopped.String(), msgE.Status)
case <-time.After(2 * interval):
t.Fatal("Timeout waiting for VM stopped message")
}
}
+1 -1
View File
@@ -12,7 +12,7 @@ import (
const VsockConfigPort uint32 = 9999
func (v *qemuVM) SendAgentConfig(ac agent.Computation) error {
conn, err := vsock.Dial(uint32(v.config.GuestCID), VsockConfigPort, nil)
conn, err := vsock.Dial(uint32(v.vmi.Config.GuestCID), VsockConfigPort, nil)
if err != nil {
return err
}
+140 -130
View File
@@ -5,29 +5,38 @@ package manager
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"log/slog"
"net"
"os"
"os/exec"
"regexp"
"strconv"
"sync"
"syscall"
"github.com/absmach/magistrala/pkg/errors"
"github.com/cenkalti/backoff/v4"
"github.com/ultravioletrs/cocos/agent"
"github.com/google/go-sev-guest/proto/check"
"github.com/google/uuid"
"github.com/ultravioletrs/cocos/manager/qemu"
"github.com/ultravioletrs/cocos/manager/vm"
"github.com/ultravioletrs/cocos/pkg/manager"
"golang.org/x/crypto/sha3"
"google.golang.org/protobuf/types/known/timestamppb"
"google.golang.org/protobuf/encoding/protojson"
)
const (
hashLength = 32
persistenceDir = "/tmp/cocos"
hashLength = 32
persistenceDir = "/tmp/cocos"
agentLogLevelKey = "AGENT_LOG_LEVEL"
agentCvmGrpcUrlKey = "AGENT_CVM_GRPC_URL"
agentCvmClientCertKey = "AGENT_CVM_GRPC_CLIENT_CERT"
agentCvmClientKey = "AGENT_CVM_GRPC_CLIENT_KEY"
agentCvmServerCaCertKey = "AGENT_CVM_GRPC_SERVER_CA_CERTS"
defClientCertPath = "/etc/certs/cert.pem"
defClientKeyPath = "/etc/certs/key.pem"
defServerCaCertPath = "/etc/certs/ca.pem"
cvmEnvironmentFile = "environment"
)
var (
@@ -45,33 +54,38 @@ var (
// ErrFailedToAllocatePort indicates no free port was found on host.
ErrFailedToAllocatePort = errors.New("failed to allocate free port on host")
errInvalidHashLength = errors.New("hash must be of byte length 32")
// ErrFailedToCalculateHash indicates that agent computation returned an error while calculating the hash of the computation.
ErrFailedToCalculateHash = errors.New("error while calculating the hash of the computation")
// ErrFailedToCreateAttestationPolicy indicates that the script to create the attestation policy failed to execute.
ErrFailedToCreateAttestationPolicy = errors.New("error while creating attestation policy")
// ErrFailedToReadPolicy indicates that the file for attestation policy could not be opened.
ErrFailedToReadPolicy = errors.New("error while opening file attestation policy")
// ErrUnmarshalFailed indicates that the file for the attestation policy could not be unmarshaled.
ErrUnmarshalFailed = errors.New("error while unmarshaling the attestation policy")
)
// Service specifies an API that must be fulfilled by the domain service
// implementation, and all of its decorators (e.g. logging & metrics).
type Service interface {
// Run create a computation.
Run(ctx context.Context, c *ComputationRunReq) (string, error)
CreateVM(ctx context.Context, req *CreateReq) (string, string, error)
// Stop stops a computation.
Stop(ctx context.Context, computationID string) error
RemoveVM(ctx context.Context, computationID string) error
// FetchAttestationPolicy measures and fetches the attestation policy.
FetchAttestationPolicy(ctx context.Context, computationID string) ([]byte, error)
// ReportBrokenConnection reports a broken connection.
ReportBrokenConnection(addr string)
// ReturnSVMInfo returns SVM information needed for attestation verification and validation.
ReturnSVMInfo(ctx context.Context) (string, int, string, string)
}
type managerService struct {
mu sync.Mutex
ap sync.Mutex
qemuCfg qemu.Config
attestationPolicyBinaryPath string
logger *slog.Logger
eventsChan chan *ClientStreamMessage
vms map[string]vm.VM
vmFactory vm.Provider
portRangeMin int
@@ -83,7 +97,7 @@ type managerService struct {
var _ Service = (*managerService)(nil)
// New instantiates the manager service implementation.
func New(cfg qemu.Config, attestationPolicyBinPath string, logger *slog.Logger, eventsChan chan *ClientStreamMessage, vmFactory vm.Provider, eosVersion string) (Service, error) {
func New(cfg qemu.Config, attestationPolicyBinPath string, logger *slog.Logger, vmFactory vm.Provider, eosVersion string) (Service, error) {
start, end, err := decodeRange(cfg.HostFwdRange)
if err != nil {
return nil, err
@@ -98,7 +112,6 @@ func New(cfg qemu.Config, attestationPolicyBinPath string, logger *slog.Logger,
qemuCfg: cfg,
logger: logger,
vms: make(map[string]vm.VM),
eventsChan: eventsChan,
vmFactory: vmFactory,
attestationPolicyBinaryPath: attestationPolicyBinPath,
portRangeMin: start,
@@ -114,52 +127,60 @@ func New(cfg qemu.Config, attestationPolicyBinPath string, logger *slog.Logger,
return ms, nil
}
func (ms *managerService) Run(ctx context.Context, c *ComputationRunReq) (string, error) {
func (ms *managerService) CreateVM(ctx context.Context, req *CreateReq) (string, string, error) {
id := uuid.New().String()
ms.mu.Lock()
cfg := ms.qemuCfg
cfg := qemu.VMInfo{
Config: ms.qemuCfg,
LaunchTCB: 0,
}
ms.mu.Unlock()
ms.publishEvent(manager.VmProvision.String(), c.Id, manager.Starting.String(), json.RawMessage{})
ac := agent.Computation{
ID: c.Id,
Name: c.Name,
Description: c.Description,
AgentConfig: agent.AgentConfig{
Port: c.AgentConfig.Port,
Host: c.AgentConfig.Host,
KeyFile: c.AgentConfig.KeyFile,
CertFile: c.AgentConfig.CertFile,
ServerCAFile: c.AgentConfig.ServerCaFile,
ClientCAFile: c.AgentConfig.ClientCaFile,
LogLevel: c.AgentConfig.LogLevel,
AttestedTls: c.AgentConfig.AttestedTls,
},
}
if len(c.Algorithm.Hash) != hashLength {
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Failed.String(), json.RawMessage{})
return "", errInvalidHashLength
tmpCertsDir, err := tempCertMount(id, req)
if err != nil {
return "", id, err
}
ac.Algorithm = agent.Algorithm{Hash: [hashLength]byte(c.Algorithm.Hash), UserKey: c.Algorithm.UserKey}
tmpEnvDir, err := tmpEnvironment(id, req)
if err != nil {
return "", id, err
}
for _, data := range c.Datasets {
if len(data.Hash) != hashLength {
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Failed.String(), json.RawMessage{})
return "", errInvalidHashLength
cfg.Config.CertsMount = tmpCertsDir
cfg.Config.EnvMount = tmpEnvDir
if ms.qemuCfg.EnableSEVSNP || ms.qemuCfg.EnableSEV {
cmd := exec.Command("sudo", fmt.Sprintf("%s/attestation_policy", ms.attestationPolicyBinaryPath), "--policy", "196608")
ms.ap.Lock()
_, err := cmd.Output()
ms.ap.Unlock()
if err != nil {
return "", id, errors.Wrap(ErrFailedToCreateAttestationPolicy, err)
}
ac.Datasets = append(ac.Datasets, agent.Dataset{Hash: [hashLength]byte(data.Hash), UserKey: data.UserKey, Filename: data.Filename})
}
for _, rc := range c.ResultConsumers {
ac.ResultConsumers = append(ac.ResultConsumers, agent.ResultConsumer{UserKey: rc.UserKey})
ms.ap.Lock()
f, err := os.ReadFile("./attestation_policy.json")
ms.ap.Unlock()
if err != nil {
return "", id, errors.Wrap(ErrFailedToReadPolicy, err)
}
var attestationPolicy check.Config
if err = protojson.Unmarshal(f, &attestationPolicy); err != nil {
return "", id, errors.Wrap(ErrUnmarshalFailed, err)
}
// Define the TCB that was present at launch of the VM.
cfg.LaunchTCB = attestationPolicy.Policy.MinimumLaunchTcb
}
agentPort, err := getFreePort(ms.portRangeMin, ms.portRangeMax)
if err != nil {
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Failed.String(), json.RawMessage{})
return "", errors.Wrap(ErrFailedToAllocatePort, err)
return "", id, errors.Wrap(ErrFailedToAllocatePort, err)
}
cfg.HostFwdAgent = agentPort
cfg.Config.HostFwdAgent = agentPort
var cid int = qemu.BaseGuestCID
for {
@@ -175,66 +196,50 @@ func (ms *managerService) Run(ctx context.Context, c *ComputationRunReq) (string
}
cid++
}
cfg.VSockConfig.GuestCID = cid
cfg.Config.VSockConfig.GuestCID = cid
ch, err := computationHash(ac)
if err != nil {
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Failed.String(), json.RawMessage{})
return "", errors.Wrap(ErrFailedToCalculateHash, err)
if cfg.Config.EnableSEVSNP {
todo := sha3.Sum256([]byte("TODO"))
// Define host-data value of QEMU for SEV-SNP, with a base64 encoding of the computation hash.
cfg.Config.SevConfig.HostData = base64.StdEncoding.EncodeToString(todo[:])
}
// Define host-data value of QEMU for SEV-SNP, with a base64 encoding of the computation hash.
cfg.SevConfig.HostData = base64.StdEncoding.EncodeToString(ch[:])
cvm := ms.vmFactory(cfg, ms.eventsLogsSender, c.Id)
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.InProgress.String(), json.RawMessage{})
cvm := ms.vmFactory(cfg, id)
if err = cvm.Start(); err != nil {
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Failed.String(), json.RawMessage{})
return "", err
return "", id, err
}
ms.mu.Lock()
ms.vms[c.Id] = cvm
ms.vms[id] = cvm
ms.mu.Unlock()
pid := cvm.GetProcess()
state := qemu.VMState{
ID: c.Id,
Config: cfg,
ID: id,
VMinfo: cfg,
PID: pid,
}
if err := ms.persistence.SaveVM(state); err != nil {
ms.logger.Error("Failed to persist VM state", "error", err)
}
err = backoff.Retry(func() error {
return cvm.SendAgentConfig(ac)
}, backoff.NewExponentialBackOff())
if err != nil {
return "", err
}
ms.mu.Lock()
if err := ms.vms[c.Id].Transition(manager.VmRunning); err != nil {
ms.logger.Warn("Failed to transition VM state", "computation", c.Id, "error", err)
if err := ms.vms[id].Transition(manager.VmRunning); err != nil {
ms.logger.Warn("Failed to transition VM state", "cvm", id, "error", err)
}
ms.mu.Unlock()
ms.publishEvent(manager.VmProvision.String(), c.Id, agent.Completed.String(), json.RawMessage{})
return fmt.Sprint(agentPort), nil
return fmt.Sprint(agentPort), id, nil
}
func (ms *managerService) Stop(ctx context.Context, computationID string) error {
func (ms *managerService) RemoveVM(ctx context.Context, computationID string) error {
ms.mu.Lock()
defer ms.mu.Unlock()
cvm, ok := ms.vms[computationID]
if !ok {
defer ms.publishEvent(manager.StopComputationRun.String(), computationID, agent.Failed.String(), json.RawMessage{})
return ErrNotFound
}
if err := cvm.Stop(); err != nil {
defer ms.publishEvent(manager.StopComputationRun.String(), computationID, agent.Failed.String(), json.RawMessage{})
return err
}
delete(ms.vms, computationID)
@@ -243,7 +248,6 @@ func (ms *managerService) Stop(ctx context.Context, computationID string) error
ms.logger.Error("Failed to delete persisted VM state", "error", err)
}
defer ms.publishEvent(manager.StopComputationRun.String(), computationID, agent.Completed.String(), json.RawMessage{})
return nil
}
@@ -295,30 +299,6 @@ func checkPortisFree(port int) bool {
return true
}
func (ms *managerService) publishEvent(event, cmpID, status string, details json.RawMessage) {
ms.eventsChan <- &ClientStreamMessage{
Message: &ClientStreamMessage_AgentEvent{
AgentEvent: &AgentEvent{
EventType: event,
ComputationId: cmpID,
Status: status,
Details: details,
Timestamp: timestamppb.Now(),
Originator: "manager",
},
},
}
}
func computationHash(ac agent.Computation) ([32]byte, error) {
jsonData, err := json.Marshal(ac)
if err != nil {
return [32]byte{}, err
}
return sha3.Sum256(jsonData), nil
}
func decodeRange(input string) (int, int, error) {
re := regexp.MustCompile(`(\d+)-(\d+)`)
matches := re.FindStringSubmatch(input)
@@ -358,7 +338,7 @@ func (ms *managerService) restoreVMs() error {
continue
}
cvm := ms.vmFactory(state.Config, ms.eventsLogsSender, state.ID)
cvm := ms.vmFactory(state.VMinfo, state.ID)
if err = cvm.SetProcess(state.PID); err != nil {
ms.logger.Warn("Failed to reattach to process", "computation", state.ID, "pid", state.PID, "error", err)
@@ -392,32 +372,62 @@ func (ms *managerService) processExists(pid int) bool {
return false
}
func (ms *managerService) eventsLogsSender(e interface{}) error {
switch msg := e.(type) {
case *vm.Event:
ms.eventsChan <- &ClientStreamMessage{
Message: &ClientStreamMessage_AgentEvent{
AgentEvent: &AgentEvent{
EventType: msg.EventType,
Timestamp: msg.Timestamp,
ComputationId: msg.ComputationId,
Originator: msg.Originator,
Status: msg.Status,
Details: msg.Details,
},
},
}
case *vm.Log:
ms.eventsChan <- &ClientStreamMessage{
Message: &ClientStreamMessage_AgentLog{
AgentLog: &AgentLog{
ComputationId: msg.ComputationId,
Level: msg.Level,
Timestamp: msg.Timestamp,
Message: msg.Message,
},
},
func tempCertMount(id string, req *CreateReq) (string, error) {
dir, err := os.MkdirTemp("/tmp", id)
if err != nil {
return "", err
}
if err = os.WriteFile(fmt.Sprintf("%s/%s", dir, "cert.pem"), req.AgentCvmClientCert, 0o644); err != nil {
return "", err
}
if err = os.WriteFile(fmt.Sprintf("%s/%s", dir, "key.pem"), req.AgentCvmClientKey, 0o644); err != nil {
return "", err
}
if err = os.WriteFile(fmt.Sprintf("%s/%s", dir, "ca.pem"), req.AgentCvmServerCaCert, 0o644); err != nil {
return "", err
}
return dir, nil
}
func tmpEnvironment(id string, req *CreateReq) (string, error) {
dir, err := os.MkdirTemp("/tmp", id)
if err != nil {
return "", err
}
envMap := map[string]string{
agentLogLevelKey: req.AgentLogLevel,
agentCvmGrpcUrlKey: req.AgentCvmServerUrl,
}
if req.AgentCvmClientCert != nil {
envMap[agentCvmClientCertKey] = defClientCertPath
}
if req.AgentCvmClientKey != nil {
envMap[agentCvmClientKey] = defClientKeyPath
}
if req.AgentCvmServerCaCert != nil {
envMap[agentCvmServerCaCertKey] = defServerCaCertPath
}
envFile, err := os.OpenFile(fmt.Sprintf("%s/%s", dir, cvmEnvironmentFile), os.O_CREATE|os.O_WRONLY, 0o644)
if err != nil {
return "", err
}
for k, v := range envMap {
if _, err = envFile.WriteString(fmt.Sprintf("%s=%s\n", k, v)); err != nil {
return "", err
}
}
return nil
if err = envFile.Close(); err != nil {
return "", err
}
return dir, nil
}
+30 -156
View File
@@ -4,7 +4,6 @@ package manager
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net"
@@ -17,7 +16,6 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"
"github.com/ultravioletrs/cocos/agent"
"github.com/ultravioletrs/cocos/manager/qemu"
persistenceMocks "github.com/ultravioletrs/cocos/manager/qemu/mocks"
"github.com/ultravioletrs/cocos/manager/vm"
@@ -29,10 +27,9 @@ func TestNew(t *testing.T) {
HostFwdRange: "6000-6100",
}
logger := slog.Default()
eventsChan := make(chan *ClientStreamMessage)
vmf := new(mocks.Provider)
service, err := New(cfg, "", logger, eventsChan, vmf.Execute, "")
service, err := New(cfg, "", logger, vmf.Execute, "")
require.NoError(t, err)
assert.NotNil(t, service)
@@ -45,67 +42,28 @@ func TestRun(t *testing.T) {
persistence := new(persistenceMocks.Persistence)
vmf.On("Execute", mock.Anything, mock.Anything, mock.Anything).Return(vmMock)
tests := []struct {
name string
req *ComputationRunReq
vmStartError error
expectedError error
name string
binaryBehavior string
vmStartError error
expectedError error
}{
{
name: "Successful run",
req: &ComputationRunReq{
Id: "test-computation",
Name: "Test Computation",
Algorithm: &Algorithm{
Hash: make([]byte, hashLength),
},
AgentConfig: &AgentConfig{},
},
vmStartError: nil,
expectedError: nil,
name: "Successful run",
binaryBehavior: "success",
vmStartError: nil,
expectedError: nil,
},
{
name: "VM start failure",
req: &ComputationRunReq{
Id: "test-computation",
Name: "Test Computation",
Algorithm: &Algorithm{
Hash: make([]byte, hashLength),
},
AgentConfig: &AgentConfig{},
},
vmStartError: assert.AnError,
expectedError: assert.AnError,
name: "VM start failure",
binaryBehavior: "success",
vmStartError: assert.AnError,
expectedError: assert.AnError,
},
{
name: "Invalid algorithm hash",
req: &ComputationRunReq{
Id: "test-computation",
Name: "Test Computation",
Algorithm: &Algorithm{
Hash: make([]byte, hashLength-1),
},
AgentConfig: &AgentConfig{},
},
vmStartError: nil,
expectedError: errInvalidHashLength,
},
{
name: "Invalid dataset hash",
req: &ComputationRunReq{
Id: "test-computation",
Name: "Test Computation",
Algorithm: &Algorithm{
Hash: make([]byte, hashLength),
},
AgentConfig: &AgentConfig{},
Datasets: []*Dataset{
{
Hash: make([]byte, hashLength-1),
},
},
},
vmStartError: nil,
expectedError: errInvalidHashLength,
name: "Invalid attestation policy",
binaryBehavior: "fail",
vmStartError: nil,
expectedError: ErrFailedToCreateAttestationPolicy,
},
}
@@ -124,29 +82,32 @@ func TestRun(t *testing.T) {
persistence.On("SaveVM", mock.Anything).Return(nil)
qemuCfg := qemu.Config{
EnableSEVSNP: true,
VSockConfig: qemu.VSockConfig{
GuestCID: 3,
},
}
logger := slog.Default()
eventsChan := make(chan *ClientStreamMessage, 10)
tempDir := CreateDummyAttestationPolicyBinary(t, tt.binaryBehavior)
defer os.RemoveAll(tempDir)
ms := &managerService{
qemuCfg: qemuCfg,
logger: logger,
vms: make(map[string]vm.VM),
eventsChan: eventsChan,
vmFactory: vmf.Execute,
persistence: persistence,
qemuCfg: qemuCfg,
attestationPolicyBinaryPath: tempDir,
logger: logger,
vms: make(map[string]vm.VM),
vmFactory: vmf.Execute,
persistence: persistence,
}
ctx := context.Background()
port, err := ms.Run(ctx, tt.req)
port, _, err := ms.CreateVM(ctx, &CreateReq{})
if tt.expectedError != nil {
assert.Error(t, err)
assert.ErrorIs(t, err, tt.expectedError)
assert.Contains(t, err.Error(), tt.expectedError.Error())
assert.Empty(t, port)
} else {
assert.NoError(t, err)
@@ -155,10 +116,6 @@ func TestRun(t *testing.T) {
}
vmf.AssertExpectations(t)
for len(eventsChan) > 0 {
<-eventsChan
}
})
}
}
@@ -202,11 +159,9 @@ func TestStop(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
logger := slog.Default()
eventsChan := make(chan *ClientStreamMessage, 10)
ms := &managerService{
logger: logger,
vms: make(map[string]vm.VM),
eventsChan: eventsChan,
persistence: persistence,
}
vmMock := new(mocks.VM)
@@ -223,7 +178,7 @@ func TestStop(t *testing.T) {
ms.vms[tt.computationID] = vmMock
}
err := ms.Stop(context.Background(), tt.computationID)
err := ms.RemoveVM(context.Background(), tt.computationID)
if tt.expectedError != nil {
assert.Error(t, err)
@@ -232,10 +187,6 @@ func TestStop(t *testing.T) {
assert.NoError(t, err)
assert.Len(t, ms.vms, 0)
}
for len(eventsChan) > 0 {
<-eventsChan
}
})
}
}
@@ -254,82 +205,6 @@ func TestGetFreePort(t *testing.T) {
assert.Greater(t, port, 6000)
}
func TestPublishEvent(t *testing.T) {
tests := []struct {
name string
event string
computationID string
status string
details json.RawMessage
}{
{
name: "Standard event",
event: "test-event",
computationID: "test-computation",
status: "test-status",
details: nil,
},
{
name: "Event with details",
event: "detailed-event",
computationID: "detailed-computation",
status: "detailed-status",
details: json.RawMessage(`{"key": "value"}`),
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
eventsChan := make(chan *ClientStreamMessage, 1)
ms := &managerService{
eventsChan: eventsChan,
}
ms.publishEvent(tt.event, tt.computationID, tt.status, tt.details)
assert.Len(t, eventsChan, 1)
event := <-eventsChan
assert.Equal(t, tt.event, event.GetAgentEvent().EventType)
assert.Equal(t, tt.computationID, event.GetAgentEvent().ComputationId)
assert.Equal(t, tt.status, event.GetAgentEvent().Status)
assert.Equal(t, "manager", event.GetAgentEvent().Originator)
assert.Equal(t, tt.details, json.RawMessage(event.GetAgentEvent().Details))
})
}
}
func TestComputationHash(t *testing.T) {
tests := []struct {
name string
computation agent.Computation
wantErr bool
}{
{
name: "Valid computation",
computation: agent.Computation{
ID: "test-id",
Name: "test-name",
},
wantErr: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
hash, err := computationHash(tt.computation)
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
assert.NotEmpty(t, hash)
hash2, _ := computationHash(tt.computation)
assert.Equal(t, hash, hash2)
}
})
}
}
func TestDecodeRange(t *testing.T) {
tests := []struct {
name string
@@ -369,7 +244,6 @@ func TestRestoreVMs(t *testing.T) {
ms := &managerService{
persistence: mockPersistence,
vms: make(map[string]vm.VM),
eventsChan: make(chan *ClientStreamMessage, 10),
vmFactory: vmf.Execute,
logger: mglog.NewMock(),
}
-131
View File
@@ -1,131 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package manager_test
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"log/slog"
"net"
"os"
"testing"
"time"
mglog "github.com/absmach/magistrala/logger"
"github.com/ultravioletrs/cocos/manager"
managergrpc "github.com/ultravioletrs/cocos/manager/api/grpc"
"golang.org/x/crypto/sha3"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/test/bufconn"
)
const (
bufSize = 1024 * 1024
keyBitSize = 4096
)
var (
lis *bufconn.Listener
algoPath = "../test/manual/algo/lin_reg.py"
dataPath = "../test/manual/data/iris.csv"
attestedTLS = false
)
type svc struct {
logger *slog.Logger
t *testing.T
}
func TestMain(m *testing.M) {
logger := mglog.NewMock()
lis = bufconn.Listen(bufSize)
s := grpc.NewServer()
manager.RegisterManagerServiceServer(s, managergrpc.NewServer(make(chan *manager.ClientStreamMessage, 1), &svc{logger: logger}))
go func() {
if err := s.Serve(lis); err != nil {
panic(err)
}
}()
code := m.Run()
s.Stop()
lis.Close()
os.Exit(code)
}
func bufDialer(context.Context, string) (net.Conn, error) {
return lis.Dial()
}
func (s *svc) Run(ctx context.Context, ipAddress string, sendMessage managergrpc.SendFunc, authInfo credentials.AuthInfo) {
privKey, err := rsa.GenerateKey(rand.Reader, keyBitSize)
if err != nil {
s.t.Fatalf("Error generating public key: %v", err)
}
pubKey, err := x509.MarshalPKIXPublicKey(&privKey.PublicKey)
if err != nil {
s.t.Fatalf("Error marshalling public key: %v", err)
}
pubPemBytes := pem.EncodeToMemory(&pem.Block{
Type: "PUBLIC KEY",
Bytes: pubKey,
})
go func() {
time.Sleep(time.Millisecond * 100)
if err := sendMessage(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_TerminateReq{
TerminateReq: &manager.Terminate{Message: "test terminate"},
},
}); err != nil {
s.t.Fatalf("failed to send terminate request: %s", err)
}
}()
go func() {
time.Sleep(time.Millisecond * 100)
algo, err := os.ReadFile(algoPath)
if err != nil {
s.t.Fatalf("failed to read algorithm file: %s", err)
return
}
data, err := os.ReadFile(dataPath)
if err != nil {
s.t.Fatalf("failed to read data file: %s", err)
return
}
pubPem, _ := pem.Decode(pubPemBytes)
algoHash := sha3.Sum256(algo)
dataHash := sha3.Sum256(data)
if err := sendMessage(&manager.ServerStreamMessage{
Message: &manager.ServerStreamMessage_RunReq{
RunReq: &manager.ComputationRunReq{
Id: "1",
Name: "sample computation",
Description: "sample description",
Datasets: []*manager.Dataset{{Hash: dataHash[:], UserKey: pubPem.Bytes}},
Algorithm: &manager.Algorithm{Hash: algoHash[:], UserKey: pubPem.Bytes},
ResultConsumers: []*manager.ResultConsumer{{UserKey: pubPem.Bytes}},
AgentConfig: &manager.AgentConfig{
Port: "7002",
LogLevel: "debug",
AttestedTls: attestedTLS,
},
},
},
}); err != nil {
s.t.Fatalf("failed to send run request: %s", err)
}
}()
}
+4 -8
View File
@@ -21,18 +21,18 @@ func New(svc manager.Service, tracer trace.Tracer) manager.Service {
return &tracingMiddleware{tracer, svc}
}
func (tm *tracingMiddleware) Run(ctx context.Context, mc *manager.ComputationRunReq) (string, error) {
func (tm *tracingMiddleware) CreateVM(ctx context.Context, req *manager.CreateReq) (string, string, error) {
ctx, span := tm.tracer.Start(ctx, "run")
defer span.End()
return tm.svc.Run(ctx, mc)
return tm.svc.CreateVM(ctx, req)
}
func (tm *tracingMiddleware) Stop(ctx context.Context, computationID string) error {
func (tm *tracingMiddleware) RemoveVM(ctx context.Context, id string) error {
ctx, span := tm.tracer.Start(ctx, "stop")
defer span.End()
return tm.svc.Stop(ctx, computationID)
return tm.svc.RemoveVM(ctx, id)
}
func (tm *tracingMiddleware) FetchAttestationPolicy(ctx context.Context, computationId string) ([]byte, error) {
@@ -42,10 +42,6 @@ func (tm *tracingMiddleware) FetchAttestationPolicy(ctx context.Context, computa
return tm.svc.FetchAttestationPolicy(ctx, computationId)
}
func (tm *tracingMiddleware) ReportBrokenConnection(addr string) {
tm.svc.ReportBrokenConnection(addr)
}
func (tm *tracingMiddleware) ReturnSVMInfo(ctx context.Context) (string, int, string, string) {
_, span := tm.tracer.Start(ctx, "return_svm_info")
defer span.End()
-108
View File
@@ -1,108 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package vm
import (
"bytes"
"io"
"log/slog"
"strings"
pkgmanager "github.com/ultravioletrs/cocos/pkg/manager"
"google.golang.org/protobuf/types/known/timestamppb"
)
var (
_ io.Writer = &Stdout{}
_ io.Writer = &Stderr{}
)
const bufSize = 1024
type Stdout struct {
EventSender EventSender
ComputationId string
}
// Write implements io.Writer.
func (s *Stdout) Write(p []byte) (n int, err error) {
inBuf := bytes.NewBuffer(p)
buf := make([]byte, bufSize)
for {
n, err := inBuf.Read(buf)
if err != nil {
if err == io.EOF {
break
}
return len(p) - inBuf.Len(), err
}
if err := sendLog(s.EventSender, s.ComputationId, string(buf[:n]), slog.LevelDebug.String()); err != nil {
return len(p) - inBuf.Len(), err
}
}
return len(p), nil
}
type Stderr struct {
EventSender EventSender
ComputationId string
StateMachine StateMachine
}
// Write implements io.Writer.
func (s *Stderr) Write(p []byte) (n int, err error) {
inBuf := bytes.NewBuffer(p)
buf := make([]byte, bufSize)
for {
n, err := inBuf.Read(buf)
if err != nil {
if err == io.EOF {
break
}
return len(p) - inBuf.Len(), err
}
if err := sendLog(s.EventSender, s.ComputationId, string(buf[:n]), ""); err != nil {
return len(p) - inBuf.Len(), err
}
}
eventMsg := &Event{
ComputationId: s.ComputationId,
EventType: s.StateMachine.State(),
Timestamp: timestamppb.Now(),
Originator: "manager",
Status: pkgmanager.Warning.String(),
}
return len(p), s.EventSender(eventMsg)
}
func sendLog(eventSender EventSender, computationID, message, level string) error {
if len(message) < 3 {
return nil
}
if level == "" {
if strings.Contains(strings.ToLower(message), "warning") {
level = slog.LevelWarn.String()
} else {
level = slog.LevelError.String()
}
}
msg := Log{
Message: message,
ComputationId: computationID,
Level: level,
Timestamp: timestamppb.Now(),
}
return eventSender(&msg)
}
-180
View File
@@ -1,180 +0,0 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package vm
import (
"log/slog"
"testing"
"time"
"github.com/stretchr/testify/assert"
pkgmanager "github.com/ultravioletrs/cocos/pkg/manager"
)
func TestStdoutWrite(t *testing.T) {
tests := []struct {
name string
input string
expectedWrites int
}{
{
name: "Single write within buffer size",
input: "Hello, World!",
expectedWrites: 1,
},
{
name: "Multiple writes within buffer size",
input: "This is a longer message that will be split into multiple writes.",
expectedWrites: 1,
},
{
name: "Large write exceeding buffer size",
input: string(make([]byte, bufSize*2+3)),
expectedWrites: 3,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
eventLogChan := make(chan interface{}, 10)
s := &Stdout{
EventSender: func(event interface{}) error {
eventLogChan <- event
return nil
},
ComputationId: "test-computation",
}
n, err := s.Write([]byte(tt.input))
assert.NoError(t, err)
assert.Equal(t, len(tt.input), n)
var receivedWrites int
for i := 0; i < tt.expectedWrites; i++ {
select {
case msg := <-eventLogChan:
receivedWrites++
agentLog := msg.(*Log)
assert.NotNil(t, agentLog)
assert.Equal(t, "test-computation", agentLog.ComputationId)
assert.Equal(t, slog.LevelDebug.String(), agentLog.Level)
assert.NotEmpty(t, agentLog.Message)
assert.NotNil(t, agentLog.Timestamp)
case <-time.After(time.Second):
t.Fatal("Timed out waiting for log message")
}
}
assert.Equal(t, tt.expectedWrites, receivedWrites)
})
}
}
func TestStderrWrite(t *testing.T) {
tests := []struct {
name string
input string
expectedWrites int
}{
{
name: "Single write within buffer size",
input: "Error: Something went wrong",
expectedWrites: 1,
},
{
name: "Multiple writes within buffer size",
input: "This is a longer error message that will be split into multiple writes.",
expectedWrites: 1,
},
{
name: "Large write exceeding buffer size",
input: string(make([]byte, bufSize*2)),
expectedWrites: 3,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
eventLogChan := make(chan interface{}, 10)
s := &Stderr{
EventSender: func(event interface{}) error {
eventLogChan <- event
return nil
},
ComputationId: "test-computation",
StateMachine: NewStateMachine(),
}
err := s.StateMachine.Transition(pkgmanager.VmRunning)
assert.NoError(t, err)
n, err := s.Write([]byte(tt.input))
assert.NoError(t, err)
assert.Equal(t, len(tt.input), n)
var receivedWrites int
for i := 0; i < tt.expectedWrites; i++ {
select {
case msg := <-eventLogChan:
receivedWrites++
switch logEv := msg.(type) {
case *Log:
assert.NotNil(t, logEv)
assert.Equal(t, "test-computation", logEv.ComputationId)
assert.Equal(t, slog.LevelError.String(), logEv.Level)
assert.NotEmpty(t, logEv.Message)
assert.NotNil(t, logEv.Timestamp)
case *Event:
assert.NotNil(t, logEv)
assert.Equal(t, "test-computation", logEv.ComputationId)
assert.Equal(t, pkgmanager.VmRunning.String(), logEv.EventType)
assert.Equal(t, pkgmanager.Warning.String(), logEv.Status)
assert.NotNil(t, logEv.Timestamp)
}
case <-time.After(time.Second):
t.Fatal("Timed out waiting for log message")
}
}
assert.Equal(t, tt.expectedWrites, receivedWrites)
})
}
}
func TestStdoutWriteErrorHandling(t *testing.T) {
eventLogChan := make(chan interface{}, 10)
s := &Stdout{
EventSender: func(event interface{}) error {
eventLogChan <- event
return assert.AnError
},
ComputationId: "test-computation",
}
message := []byte("This should fail")
n, err := s.Write(message)
assert.Error(t, err)
assert.Equal(t, len(message), n)
assert.Equal(t, assert.AnError, err)
}
func TestStderrWriteErrorHandling(t *testing.T) {
eventLogChan := make(chan interface{}, 10)
s := &Stderr{
EventSender: func(event interface{}) error {
eventLogChan <- event
return assert.AnError
},
ComputationId: "test-computation",
}
message := []byte("This should fail")
n, err := s.Write(message)
assert.Error(t, err)
assert.Equal(t, len(message), n)
assert.Equal(t, assert.AnError, err)
}
+10 -11
View File
@@ -23,17 +23,17 @@ func (_m *Provider) EXPECT() *Provider_Expecter {
return &Provider_Expecter{mock: &_m.Mock}
}
// Execute provides a mock function with given fields: config, eventSender, computationId
func (_m *Provider) Execute(config interface{}, eventSender vm.EventSender, computationId string) vm.VM {
ret := _m.Called(config, eventSender, computationId)
// Execute provides a mock function with given fields: config, computationId
func (_m *Provider) Execute(config interface{}, computationId string) vm.VM {
ret := _m.Called(config, computationId)
if len(ret) == 0 {
panic("no return value specified for Execute")
}
var r0 vm.VM
if rf, ok := ret.Get(0).(func(interface{}, vm.EventSender, string) vm.VM); ok {
r0 = rf(config, eventSender, computationId)
if rf, ok := ret.Get(0).(func(interface{}, string) vm.VM); ok {
r0 = rf(config, computationId)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(vm.VM)
@@ -50,15 +50,14 @@ type Provider_Execute_Call struct {
// Execute is a helper method to define mock.On call
// - config interface{}
// - eventSender vm.EventSender
// - computationId string
func (_e *Provider_Expecter) Execute(config interface{}, eventSender interface{}, computationId interface{}) *Provider_Execute_Call {
return &Provider_Execute_Call{Call: _e.mock.On("Execute", config, eventSender, computationId)}
func (_e *Provider_Expecter) Execute(config interface{}, computationId interface{}) *Provider_Execute_Call {
return &Provider_Execute_Call{Call: _e.mock.On("Execute", config, computationId)}
}
func (_c *Provider_Execute_Call) Run(run func(config interface{}, eventSender vm.EventSender, computationId string)) *Provider_Execute_Call {
func (_c *Provider_Execute_Call) Run(run func(config interface{}, computationId string)) *Provider_Execute_Call {
_c.Call.Run(func(args mock.Arguments) {
run(args[0].(interface{}), args[1].(vm.EventSender), args[2].(string))
run(args[0].(interface{}), args[1].(string))
})
return _c
}
@@ -68,7 +67,7 @@ func (_c *Provider_Execute_Call) Return(_a0 vm.VM) *Provider_Execute_Call {
return _c
}
func (_c *Provider_Execute_Call) RunAndReturn(run func(interface{}, vm.EventSender, string) vm.VM) *Provider_Execute_Call {
func (_c *Provider_Execute_Call) RunAndReturn(run func(interface{}, string) vm.VM) *Provider_Execute_Call {
_c.Call.Return(run)
return _c
}
+1 -3
View File
@@ -21,7 +21,7 @@ type VM interface {
GetConfig() interface{}
}
type Provider func(config interface{}, eventSender EventSender, computationId string) VM
type Provider func(config interface{}, computationId string) VM
type Event struct {
EventType string
@@ -38,5 +38,3 @@ type Log struct {
Level string
Timestamp *timestamppb.Timestamp
}
type EventSender func(event interface{}) error
+12
View File
@@ -62,6 +62,11 @@ packages:
dir: "{{.InterfaceDir}}/mocks"
filename: "service.go"
mockname: "{{.InterfaceName}}"
ManagerServiceClient:
config:
dir: "{{.InterfaceDir}}/mocks"
filename: "manager_service_client.go"
mockname: "{{.InterfaceName}}"
github.com/ultravioletrs/cocos/manager/qemu:
interfaces:
Persistence:
@@ -93,3 +98,10 @@ packages:
dir: "{{.InterfaceDir}}/mocks"
filename: "sdk.go"
mockname: "{{.InterfaceName}}"
github.com/ultravioletrs/cocos/agent/cvms/server:
interfaces:
AgentServerProvider:
config:
dir: "{{.InterfaceDir}}/mocks"
filename: "server.go"
mockname: "{{.InterfaceName}}"
+31 -19
View File
@@ -30,12 +30,13 @@ const (
)
const (
NO_ERROR = 0
ERROR_ZERO_RETURN = 6
ERROR_WANT_READ = 2
ERROR_WANT_WRITE = 3
ERROR_SYSCALL = 5
ERROR_SSL = 1
noError = 0
errorZeroReturn = 6
errorWantRead = 2
errorWantWrite = 3
errorSyscall = 5
errorSsl = 1
waitTime = 2
)
var (
@@ -228,21 +229,21 @@ func (c *ATLSConn) Read(b []byte) (int, error) {
// handle specific error codes returned by SSL_get_error.
switch errCode {
case NO_ERROR:
case noError:
return n, nil // no error.
case ERROR_ZERO_RETURN:
fmt.Fprintf(os.Stderr, "Connection closed by peer")
case errorZeroReturn:
fmt.Fprintf(os.Stdout, "Connection closed by peer")
return 0, io.EOF // connection closed.
case ERROR_WANT_READ:
case errorWantRead:
fmt.Fprintf(os.Stderr, "Operation read incomplete, retry later")
return 0, nil // non-fatal, just retry later.
case ERROR_WANT_WRITE:
case errorWantWrite:
fmt.Fprintf(os.Stderr, "Operation write incomplete, retry later")
return 0, nil // non-fatal, just retry later.
case ERROR_SYSCALL:
case errorSyscall:
fmt.Fprintf(os.Stderr, "I/O error")
return 0, syscall.ECONNRESET // return connection reset error.
case ERROR_SSL:
case errorSsl:
fmt.Fprintf(os.Stderr, "I/O error")
return 0, syscall.ECONNRESET // return connection reset error.
default:
@@ -280,13 +281,24 @@ func (c *ATLSConn) Close() error {
return nil
}
ret := C.tls_close(c.tlsConn)
for {
ret := C.tls_close(c.tlsConn)
if int(ret) < 0 {
c.tlsConn = nil
return errTLSConn
} else if int(ret) == 1 {
c.tlsConn = nil
if int(ret) == 0 {
c.fdDelayMutex.Unlock()
c.fdWriteMutex.Unlock()
c.fdReadMutex.Unlock()
time.Sleep(waitTime * time.Millisecond)
c.fdDelayMutex.Lock()
c.fdWriteMutex.Lock()
c.fdReadMutex.Lock()
} else if int(ret) < 0 {
c.tlsConn = nil
return errTLSConn
} else if int(ret) == 1 {
c.tlsConn = nil
break;
}
}
return nil
+31 -9
View File
@@ -353,16 +353,38 @@ int tls_close(tls_connection *conn) {
if (conn->ssl != NULL) {
int ret = 0;
while (ret == 0) {
ret = SSL_shutdown(conn->ssl);
if (SSL_has_pending(conn->ssl) == 1 || (SSL_get_shutdown(conn->ssl) & SSL_SENT_SHUTDOWN)) {
int num = SSL_pending(conn->ssl);
char c[num];
int res = 0;
int end = 0;
if (ret < 0) {
fprintf(stderr, "SSL did not shutdown correctly\n");
free(conn);
close(conn->socket_fd);
conn = NULL;
return -1;
res = SSL_read(conn->ssl, (void*)c, num);
res = SSL_get_error(conn->ssl, res);
if (res == SSL_ERROR_ZERO_RETURN) {
end = 1;
} else if (res != SSL_ERROR_NONE) {
fprintf(stderr, "SSL_read failed in TLS close call\n");
end = 1;
}
if ((SSL_get_shutdown(conn->ssl) & SSL_RECEIVED_SHUTDOWN) || end == 1) {
ret = SSL_shutdown(conn->ssl);
}
} else {
ret = SSL_shutdown(conn->ssl);
}
if (ret < 0) {
ret = SSL_get_error(conn->ssl, ret);
fprintf(stderr, "SSL did not shutdown correctly, error code: %d\n", ret);
free(conn);
close(conn->socket_fd);
conn = NULL;
return -1;
} else if (ret == 0) {
return 0;
}
conn->ssl = NULL;
}
@@ -381,7 +403,7 @@ int tls_close(tls_connection *conn) {
return 1;
}
return 0;
return 1;
}
char* tls_return_addr(struct sockaddr_storage *addr) {
+5 -1
View File
@@ -72,6 +72,10 @@ type ManagerClientConfig struct {
BaseConfig
}
type CVMClientConfig struct {
BaseConfig
}
func (a BaseConfig) GetBaseConfig() BaseConfig {
return a
}
@@ -80,7 +84,7 @@ func (a AgentClientConfig) GetBaseConfig() BaseConfig {
return a.BaseConfig
}
func (a ManagerClientConfig) GetBaseConfig() BaseConfig {
func (a CVMClientConfig) GetBaseConfig() BaseConfig {
return a.BaseConfig
}
+18
View File
@@ -0,0 +1,18 @@
// Copyright (c) Ultraviolet
// SPDX-License-Identifier: Apache-2.0
package cvm
import (
"github.com/ultravioletrs/cocos/agent/cvms"
"github.com/ultravioletrs/cocos/pkg/clients/grpc"
)
// NewManagerClient creates new manager gRPC client instance.
func NewCVMClient(cfg grpc.CVMClientConfig) (grpc.Client, cvms.ServiceClient, error) {
client, err := grpc.NewClient(cfg)
if err != nil {
return nil, nil, err
}
return client, cvms.NewServiceClient(client.Connection()), nil
}
+1 -1
View File
@@ -9,5 +9,5 @@ edition = "2021"
clap = { version = "4.0", features = ["derive"] }
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
sev = "4.0.0"
sev = "5.0.0"
base64 = "0.22.1"

Some files were not shown because too many files have changed in this diff Show More