Compare commits
46 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f407e4c65c | |||
| ac5e062b72 | |||
| b4e9eb8830 | |||
| 2240742945 | |||
| e6504b21d4 | |||
| 5aa69a8a6d | |||
| 54cd178092 | |||
| 32470c7f1c | |||
| 57c3c0525e | |||
| e356593f73 | |||
| 0e033831d5 | |||
| 0d721228d5 | |||
| e49b75b7d8 | |||
| eab691b1a1 | |||
| fd6494873d | |||
| 6cbfc1fee0 | |||
| b986ae132a | |||
| f24a47969e | |||
| a0bc1827d9 | |||
| f2869cea30 | |||
| 77cf447717 | |||
| 790ed3efb3 | |||
| 5ae7933d41 | |||
| 2ab977ed18 | |||
| 1eae9a34f0 | |||
| 4e7748b059 | |||
| 582f67cade | |||
| 9e23ba6b50 | |||
| 3f8a3ac0f1 | |||
| e0b55ab057 | |||
| 421f2773c7 | |||
| f717f9982f | |||
| 44dbfde0b4 | |||
| 713511902d | |||
| 80531c9c28 | |||
| 37daf2104f | |||
| 9afdd4570c | |||
| 55fbe66fe7 | |||
| 3794c97c1e | |||
| 848623766d | |||
| 4cd09ec900 | |||
| 3ed5e1e5b5 | |||
| 3f372ff7b3 | |||
| c453c41fd2 | |||
| 5c8ac61af6 | |||
| a02e1b91d9 |
@@ -0,0 +1,29 @@
|
||||
name: Badge - Compatibility
|
||||
|
||||
on:
|
||||
workflow_run:
|
||||
workflows:
|
||||
- Examples - Backward Compatibility
|
||||
types: [completed]
|
||||
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
actions: read
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
badge:
|
||||
if: ${{ github.event_name == 'workflow_dispatch' || (github.event_name == 'workflow_run' && github.event.workflow_run.head_branch == 'main') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/github-script@v8
|
||||
with:
|
||||
github-token: ${{ secrets.GITHUB_TOKEN }}
|
||||
script: |
|
||||
const badgeAggregation = require('./scripts/badge_aggregation.js');
|
||||
const dependencies = [
|
||||
{ workflow: 'examples-compat.yml', label: 'examples-compat', variants: ['legacy', 'stable'] },
|
||||
];
|
||||
await badgeAggregation({ github, context, core, dependencies });
|
||||
@@ -0,0 +1,31 @@
|
||||
name: Badge - Unit Test
|
||||
|
||||
on:
|
||||
workflow_run:
|
||||
workflows:
|
||||
- CPU Test
|
||||
- GPU Test
|
||||
types: [completed]
|
||||
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
actions: read
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
badge:
|
||||
if: ${{ github.event_name == 'workflow_dispatch' || (github.event_name == 'workflow_run' && github.event.workflow_run.head_branch == 'main') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/github-script@v8
|
||||
with:
|
||||
github-token: ${{ secrets.GITHUB_TOKEN }}
|
||||
script: |
|
||||
const badgeAggregation = require('./scripts/badge_aggregation.js');
|
||||
const dependencies = [
|
||||
{ workflow: 'tests-full.yml', label: 'tests-full', variants: ['legacy', 'stable'] },
|
||||
{ workflow: 'tests.yml', label: 'tests', variants: ['legacy', 'stable', 'Lint', 'documentation', 'JavaScript'] },
|
||||
];
|
||||
await badgeAggregation({ github, context, core, dependencies });
|
||||
@@ -0,0 +1,33 @@
|
||||
name: Dashboard
|
||||
permissions:
|
||||
contents: read
|
||||
on:
|
||||
schedule:
|
||||
# Every day at 5 AM UTC+8
|
||||
- cron: '0 21 * * *'
|
||||
|
||||
workflow_dispatch:
|
||||
|
||||
push:
|
||||
branches: [ main, stable/**/* ]
|
||||
|
||||
jobs:
|
||||
dashboard:
|
||||
name: Chromatic
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Run Chromatic
|
||||
uses: chromaui/action@v13
|
||||
with:
|
||||
projectToken: ${{ secrets.CHROMATIC_PROJECT_TOKEN }}
|
||||
workingDir: dashboard
|
||||
exitZeroOnChanges: false
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'APO - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'Calc-X - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
@@ -22,12 +22,12 @@ run-name: >-
|
||||
|| format('Calc-X - {0}', github.event_name) }}
|
||||
|
||||
jobs:
|
||||
calc-x:
|
||||
calc-x-perf:
|
||||
if: >
|
||||
github.event_name != 'repository_dispatch' ||
|
||||
github.event.action == 'ci-calc-x' ||
|
||||
github.event.action == 'ci-all'
|
||||
name: Calc-X (Python ${{ matrix.python-version }}, ${{ matrix.setup-script }})
|
||||
name: Calc-X Performance (Python ${{ matrix.python-version }}, ${{ matrix.setup-script }})
|
||||
runs-on: [self-hosted, 1ES.Pool=agl-runner-gpu]
|
||||
timeout-minutes: 90
|
||||
strategy:
|
||||
@@ -74,7 +74,7 @@ jobs:
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-calc-x-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
name: dependencies-calc-x-performance-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
|
||||
@@ -116,13 +116,11 @@ jobs:
|
||||
# Don't ask why. Don't touch this.
|
||||
- name: Calc-X training
|
||||
run: |
|
||||
set -ex
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
../../scripts/restart_ray.sh
|
||||
sleep 5
|
||||
PYTHONUNBUFFERED=1 python train_calc_agent.py --val-file data/test_mini.parquet --ci
|
||||
sleep 10
|
||||
python train_calc_agent.py --val-file data/test_mini.parquet --ci
|
||||
shell: bash
|
||||
env:
|
||||
WANDB_BASE_URL: ${{ secrets.MSR_WANDB_BASE_URL }}
|
||||
@@ -137,14 +135,118 @@ jobs:
|
||||
WANDB_BASE_URL: ${{ secrets.MSR_WANDB_BASE_URL }}
|
||||
WANDB_API_KEY: ${{ secrets.MSR_WANDB_API_KEY }}
|
||||
|
||||
- name: Calc-X training LLM Proxy
|
||||
calc-x-variants:
|
||||
if: >
|
||||
github.event_name != 'repository_dispatch' ||
|
||||
github.event.action == 'ci-calc-x' ||
|
||||
github.event.action == 'ci-all'
|
||||
name: Calc-X Variants (Python ${{ matrix.python-version }}, ${{ matrix.setup-script }})
|
||||
runs-on: [self-hosted, 1ES.Pool=agl-runner-gpu]
|
||||
timeout-minutes: 90
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- python-version: '3.10'
|
||||
setup-script: 'legacy'
|
||||
- python-version: '3.12'
|
||||
setup-script: 'stable'
|
||||
- python-version: '3.13'
|
||||
setup-script: 'latest'
|
||||
fail-fast: false
|
||||
steps:
|
||||
- name: Check GPU status
|
||||
run: nvidia-smi
|
||||
- name: Check disk space
|
||||
run: df -h
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event_name == 'repository_dispatch' && github.event.client_payload.pr_ref || (github.event.pull_request.number && format('refs/pull/{0}/merge', github.event.pull_request.number)) || github.ref }}
|
||||
- uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Upgrade dependencies (latest)
|
||||
run: uv lock --upgrade
|
||||
if: matrix.setup-script == 'latest'
|
||||
- name: Sync dependencies (latest)
|
||||
run: |
|
||||
uv sync --frozen --no-default-groups --extra verl \
|
||||
--group dev --group experiment --group agents --group torch-gpu-stable
|
||||
if: matrix.setup-script == 'latest'
|
||||
- name: Sync dependencies (stable & legacy)
|
||||
run: |
|
||||
uv sync --frozen --no-default-groups --extra verl \
|
||||
--group dev --group experiment --group agents --group torch-gpu-${{ matrix.setup-script }}
|
||||
if: matrix.setup-script != 'latest'
|
||||
- name: Freeze dependencies
|
||||
run: |
|
||||
set -ex
|
||||
uv pip freeze | tee requirements-freeze.txt
|
||||
echo "UV_LOCKED=1" >> $GITHUB_ENV
|
||||
echo "UV_NO_SYNC=1" >> $GITHUB_ENV
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-calc-x-variants-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
|
||||
- name: Launch LiteLLM Proxy
|
||||
run: |
|
||||
./scripts/litellm_run.sh
|
||||
env:
|
||||
AZURE_API_BASE: ${{ secrets.AZURE_GROUP_SUBSCRIPTION_API_BASE }}
|
||||
AZURE_API_KEY: ${{ secrets.AZURE_GROUP_SUBSCRIPTION_API_KEY }}
|
||||
|
||||
- name: Prepare Calc-X dataset
|
||||
run: |
|
||||
set -ex
|
||||
cd examples/calc_x
|
||||
uv run gdown --fuzzy https://drive.google.com/file/d/1FQMyKLLd6hP9dw9rfZn1EZOWNvKaDsqw/view
|
||||
unzip calc-x-data.zip -d data
|
||||
rm calc-x-data.zip
|
||||
|
||||
- name: Calc-X MCP sanity check
|
||||
run: |
|
||||
set -ex
|
||||
cd examples/calc_x
|
||||
uv run tests/test_mcp_calculator.py
|
||||
env:
|
||||
OPENAI_API_BASE: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
- name: Calc-X sanity check
|
||||
run: |
|
||||
set -ex
|
||||
cd examples/calc_x
|
||||
uv run legacy_calc_agent_debug.py
|
||||
env:
|
||||
OPENAI_BASE_URL: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
|
||||
- name: Training with local model
|
||||
run: |
|
||||
set -ex
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
../../scripts/restart_ray.sh
|
||||
sleep 5
|
||||
PYTHONUNBUFFERED=1 python train_calc_agent.py --val-file data/test_mini.parquet --ci --llm-proxy
|
||||
hf download Qwen/Qwen2.5-0.5B-Instruct --local-dir data/qwen_model
|
||||
PYTHONUNBUFFERED=1 python train_calc_agent.py --val-file data/test_mini.parquet --ci-fast --model $(realpath data/qwen_model)
|
||||
sleep 10
|
||||
shell: bash
|
||||
env:
|
||||
WANDB_BASE_URL: ${{ secrets.MSR_WANDB_BASE_URL }}
|
||||
WANDB_API_KEY: ${{ secrets.MSR_WANDB_API_KEY }}
|
||||
id: calc_x_train_local_model
|
||||
|
||||
- name: Training with LLM Proxy
|
||||
run: |
|
||||
set -ex
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
../../scripts/restart_ray.sh
|
||||
sleep 5
|
||||
PYTHONUNBUFFERED=1 python train_calc_agent.py --val-file data/test_mini.parquet --ci-fast --llm-proxy
|
||||
sleep 10
|
||||
shell: bash
|
||||
env:
|
||||
@@ -152,7 +254,7 @@ jobs:
|
||||
WANDB_API_KEY: ${{ secrets.MSR_WANDB_API_KEY }}
|
||||
id: calc_x_train_llm_proxy
|
||||
|
||||
- name: Calc-X training with external store
|
||||
- name: Training with external store
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
@@ -182,7 +284,7 @@ jobs:
|
||||
WANDB_API_KEY: ${{ secrets.MSR_WANDB_API_KEY }}
|
||||
id: calc_x_train_external_store
|
||||
|
||||
- name: Calc-X training with role-based environment variables
|
||||
- name: Training with role-based environment variables
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'Backward Compatibility - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
@@ -54,6 +54,10 @@ jobs:
|
||||
run: |
|
||||
uv sync --frozen --no-default-groups --extra apo --extra verl \
|
||||
--group dev --group experiment --group agents --group torch-gpu-${{ matrix.setup-script }}
|
||||
- name: Override VERL (stable)
|
||||
run: |
|
||||
uv pip install verl==0.5.0 vllm==0.10.2
|
||||
if: matrix.setup-script == 'stable'
|
||||
- name: Freeze dependencies
|
||||
run: |
|
||||
set -ex
|
||||
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'Spider - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'Unsloth - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
|
||||
@@ -26,6 +26,14 @@ jobs:
|
||||
- name: Sync dependencies
|
||||
run: uv sync --frozen --no-default-groups --group dev
|
||||
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Build dashboard
|
||||
run: cd dashboard && npm run build
|
||||
|
||||
- name: Get current version
|
||||
id: get_version
|
||||
run: |
|
||||
|
||||
@@ -60,6 +60,14 @@ jobs:
|
||||
- name: Sync dependencies
|
||||
run: uv sync --frozen --no-default-groups --group dev
|
||||
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Build dashboard
|
||||
run: cd dashboard && npm run build
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
uv build
|
||||
|
||||
@@ -14,7 +14,7 @@ on:
|
||||
run-name: >-
|
||||
${{ github.event_name == 'repository_dispatch'
|
||||
&& format(
|
||||
'PR #{0} - Label {1} - {2}',
|
||||
'GPU Test - PR #{0} - {1} - {2}',
|
||||
github.event.client_payload.pull_number,
|
||||
github.event.client_payload.ci_label,
|
||||
github.event.client_payload.correlation_id
|
||||
@@ -46,7 +46,7 @@ jobs:
|
||||
run: nvidia-smi
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.number && format('refs/pull/{0}/merge', github.event.pull_request.number) || github.ref }}
|
||||
ref: ${{ github.event_name == 'repository_dispatch' && github.event.client_payload.pr_ref || (github.event.pull_request.number && format('refs/pull/{0}/merge', github.event.pull_request.number)) || github.ref }}
|
||||
- uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
@@ -69,10 +69,18 @@ jobs:
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
name: dependencies-tests-full-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Build dashboard
|
||||
run: cd dashboard && npm run build
|
||||
|
||||
- name: Launch LiteLLM Proxy
|
||||
run: |
|
||||
./scripts/litellm_run.sh
|
||||
@@ -87,3 +95,177 @@ jobs:
|
||||
PYTEST_ADDOPTS: "--color=yes"
|
||||
OPENAI_BASE_URL: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
|
||||
minimal-examples:
|
||||
if: >
|
||||
github.event_name != 'repository_dispatch' ||
|
||||
github.event.action == 'ci-gpu' ||
|
||||
github.event.action == 'ci-all'
|
||||
name: Minimal Examples with Python ${{ matrix.python-version }} (${{ matrix.setup-script }})
|
||||
|
||||
runs-on: [self-hosted, 1ES.Pool=agl-runner-gpu]
|
||||
timeout-minutes: 30
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- python-version: '3.10'
|
||||
setup-script: 'legacy'
|
||||
- python-version: '3.12'
|
||||
setup-script: 'stable'
|
||||
- python-version: '3.13'
|
||||
setup-script: 'latest'
|
||||
fail-fast: false
|
||||
steps:
|
||||
- name: Check GPU status
|
||||
run: nvidia-smi
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event_name == 'repository_dispatch' && github.event.client_payload.pr_ref || (github.event.pull_request.number && format('refs/pull/{0}/merge', github.event.pull_request.number)) || github.ref }}
|
||||
- uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Upgrade dependencies (latest)
|
||||
run: uv lock --upgrade
|
||||
if: matrix.setup-script == 'latest'
|
||||
- name: Sync dependencies (latest)
|
||||
run: uv sync --frozen --no-default-groups --extra apo --group dev --group agents --group torch-gpu-stable
|
||||
if: matrix.setup-script == 'latest'
|
||||
- name: Sync dependencies (stable & legacy)
|
||||
run: uv sync --frozen --no-default-groups --extra apo --group dev --group agents --group torch-gpu-${{ matrix.setup-script }}
|
||||
if: matrix.setup-script != 'latest'
|
||||
- name: Freeze dependencies
|
||||
run: |
|
||||
set -ex
|
||||
uv pip freeze | tee requirements-freeze.txt
|
||||
echo "UV_LOCKED=1" >> $GITHUB_ENV
|
||||
echo "UV_NO_SYNC=1" >> $GITHUB_ENV
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-minimal-examples-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
|
||||
- name: Launch LiteLLM Proxy
|
||||
run: |
|
||||
./scripts/litellm_run.sh
|
||||
env:
|
||||
AZURE_API_BASE: ${{ secrets.AZURE_GROUP_SUBSCRIPTION_API_BASE }}
|
||||
AZURE_API_KEY: ${{ secrets.AZURE_GROUP_SUBSCRIPTION_API_KEY }}
|
||||
|
||||
- name: Write Traces via Otel Tracer
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
python write_traces.py otel
|
||||
|
||||
- name: Write Traces via AgentOps Tracer
|
||||
env:
|
||||
OPENAI_BASE_URL: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
python write_traces.py agentops
|
||||
|
||||
- name: Write Traces via Otel Tracer with Client
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
agl store --port 45993 --log-level DEBUG &
|
||||
sleep 5
|
||||
python write_traces.py otel --use-client
|
||||
pkill -f agl && echo "SIGTERM sent to agl" || echo "No agl process found"
|
||||
while pgrep -f agl; do
|
||||
echo "Waiting for agl to finish..."
|
||||
sleep 5
|
||||
done
|
||||
|
||||
- name: Write Traces via AgentOps Tracer with Client
|
||||
env:
|
||||
OPENAI_BASE_URL: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
agl store --port 45993 --log-level DEBUG &
|
||||
sleep 5
|
||||
python write_traces.py agentops --use-client
|
||||
pkill -f agl && echo "SIGTERM sent to agl" || echo "No agl process found"
|
||||
while pgrep -f agl; do
|
||||
echo "Waiting for agl to finish..."
|
||||
sleep 5
|
||||
done
|
||||
|
||||
- name: vLLM Server
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
python vllm_server.py Qwen/Qwen2.5-0.5B-Instruct
|
||||
|
||||
- name: LLM Proxy (OpenAI backend)
|
||||
env:
|
||||
OPENAI_API_BASE: http://localhost:12306/
|
||||
OPENAI_API_KEY: dummy
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
|
||||
python llm_proxy.py openai gpt-4.1-mini &
|
||||
|
||||
LLM_PROXY_READY=0
|
||||
for attempt in $(seq 1 30); do
|
||||
if curl -sSf http://localhost:43886/health > /dev/null 2>&1; then
|
||||
LLM_PROXY_READY=1
|
||||
break
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
if [[ "$LLM_PROXY_READY" != "1" ]]; then
|
||||
echo "LLM proxy failed to become healthy" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
python llm_proxy.py test gpt-4.1-mini
|
||||
|
||||
pkill -f llm_proxy.py && echo "SIGTERM sent to llm_proxy.py" || echo "No llm_proxy.py process found"
|
||||
while pgrep -f llm_proxy.py; do
|
||||
echo "Waiting for llm_proxy.py to finish..."
|
||||
sleep 5
|
||||
done
|
||||
|
||||
- name: LLM Proxy (vLLM backend)
|
||||
if: matrix.setup-script != 'legacy' # Skip if return_token_ids is not supported
|
||||
run: |
|
||||
set -euo pipefail
|
||||
source .venv/bin/activate
|
||||
cd examples/minimal
|
||||
python llm_proxy.py vllm Qwen/Qwen2.5-0.5B-Instruct &
|
||||
|
||||
LLM_PROXY_READY=0
|
||||
for attempt in $(seq 1 30); do
|
||||
if curl -sSf http://localhost:43886/health > /dev/null 2>&1; then
|
||||
LLM_PROXY_READY=1
|
||||
break
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
if [[ "$LLM_PROXY_READY" != "1" ]]; then
|
||||
echo "LLM proxy failed to become healthy" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
python llm_proxy.py test Qwen/Qwen2.5-0.5B-Instruct
|
||||
|
||||
pkill -f llm_proxy.py && echo "SIGTERM sent to llm_proxy.py" || echo "No llm_proxy.py process found"
|
||||
while pgrep -f llm_proxy.py; do
|
||||
echo "Waiting for llm_proxy.py to finish..."
|
||||
sleep 5
|
||||
done
|
||||
|
||||
@@ -45,10 +45,11 @@ jobs:
|
||||
--group agents \
|
||||
--no-default-groups
|
||||
if: matrix.setup == 'slow'
|
||||
# This pre-commit skips JavaScript on purpose.
|
||||
- name: Run pre-commit
|
||||
uses: pre-commit/action@v3.0.1
|
||||
- name: Check Python headers
|
||||
run: uv run --locked --no-sync scripts/check_python_headers.py
|
||||
run: uv run --locked --no-sync scripts/check_headers.py
|
||||
- name: Run Black
|
||||
run: uv run --locked --no-sync black --check .
|
||||
- name: Run isort
|
||||
@@ -60,6 +61,28 @@ jobs:
|
||||
run: uv run --locked --no-sync pyright -p pyrightconfig.json
|
||||
if: matrix.setup == 'slow'
|
||||
|
||||
lint-js:
|
||||
name: Lint - JavaScript
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Run ESLint
|
||||
run: cd dashboard && npm run eslint
|
||||
- name: Run Prettier
|
||||
run: cd dashboard && npm run prettier
|
||||
- name: Run Stylelint
|
||||
run: cd dashboard && npm run stylelint
|
||||
- name: Run Typecheck
|
||||
run: cd dashboard && npm run typecheck
|
||||
- name: Verify build
|
||||
run: cd dashboard && npm run build
|
||||
|
||||
docs:
|
||||
name: Build documentation
|
||||
runs-on: ubuntu-latest
|
||||
@@ -132,8 +155,39 @@ jobs:
|
||||
name: dependencies-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Build dashboard
|
||||
run: cd dashboard && npm run build
|
||||
|
||||
- name: Run tests
|
||||
run: |
|
||||
uv run pytest -v --durations=0 tests
|
||||
env:
|
||||
PYTEST_ADDOPTS: "--color=yes"
|
||||
|
||||
test-js:
|
||||
name: Test - JavaScript
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
- uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
python-version: '3.12'
|
||||
- name: Sync Python dependencies
|
||||
run: uv sync --frozen --no-default-groups --extra apo --group dev --group agents --group core-stable
|
||||
- name: Install JavaScript dependencies
|
||||
run: cd dashboard && npm ci
|
||||
- name: Run vitest
|
||||
run: cd dashboard && npm run vitest
|
||||
|
||||
@@ -207,3 +207,9 @@ cython_debug/
|
||||
|
||||
# Claude
|
||||
.claude/*.local.json
|
||||
|
||||
# Dashboard generated files
|
||||
agentlightning/dashboard/**/*.css
|
||||
agentlightning/dashboard/**/*.js
|
||||
agentlightning/dashboard/**/*.html
|
||||
agentlightning/dashboard/**/*.svg
|
||||
|
||||
@@ -24,3 +24,53 @@ repos:
|
||||
pass_filenames: false
|
||||
always_run: true
|
||||
args: ["."]
|
||||
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: prettier
|
||||
name: prettier (dashboard)
|
||||
language: system
|
||||
pass_filenames: false
|
||||
always_run: true
|
||||
entry: >
|
||||
bash -c '
|
||||
cd dashboard || exit 1
|
||||
if [ -d node_modules ]; then
|
||||
echo "✅ node_modules already exists"
|
||||
npx prettier --cache --write "**/*.{ts,tsx,mjs,cjs}"
|
||||
else
|
||||
echo "⚠️ node_modules not found — npx is not reliable. Skipping."
|
||||
fi
|
||||
'
|
||||
|
||||
- id: eslint
|
||||
name: eslint (dashboard)
|
||||
language: system
|
||||
pass_filenames: false
|
||||
always_run: true
|
||||
entry: >
|
||||
bash -c '
|
||||
cd dashboard || exit 1
|
||||
if [ -d node_modules ]; then
|
||||
echo "✅ node_modules already exists"
|
||||
npx eslint --cache --fix .
|
||||
else
|
||||
echo "⚠️ node_modules not found — npx is not reliable. Skipping."
|
||||
fi
|
||||
'
|
||||
|
||||
- id: stylelint
|
||||
name: stylelint (dashboard)
|
||||
language: system
|
||||
pass_filenames: false
|
||||
always_run: true
|
||||
entry: >
|
||||
bash -c '
|
||||
cd dashboard || exit 1
|
||||
if [ -d node_modules ]; then
|
||||
echo "✅ node_modules already exists"
|
||||
npx stylelint --cache --fix "**/*.css"
|
||||
else
|
||||
echo "⚠️ node_modules not found — npx is not reliable. Skipping."
|
||||
fi
|
||||
'
|
||||
|
||||
@@ -4,10 +4,11 @@
|
||||
|
||||
# Agent Lightning⚡
|
||||
|
||||
[](https://github.com/microsoft/agent-lightning/actions/workflows/tests-full.yml)
|
||||
[](https://github.com/microsoft/agent-lightning/actions/workflows/badge-unit.yml)
|
||||
[](https://microsoft.github.io/agent-lightning/)
|
||||
[](https://badge.fury.io/py/agentlightning)
|
||||
[](LICENSE)
|
||||
[](https://deepwiki.com/microsoft/agent-lightning)
|
||||
[](https://discord.gg/RYk7CdvDR7)
|
||||
|
||||
**The absolute trainer to light up AI agents.**
|
||||
@@ -33,12 +34,19 @@ Read more on our [documentation website](https://microsoft.github.io/agent-light
|
||||
pip install agentlightning
|
||||
```
|
||||
|
||||
For the latest nightly build (cutting-edge features), you can install from Test PyPI:
|
||||
|
||||
```bash
|
||||
pip install --upgrade --index-url https://test.pypi.org/simple/ --extra-index-url https://pypi.org/simple/ agentlightning
|
||||
```
|
||||
|
||||
Please refer to our [installation guide](https://microsoft.github.io/agent-lightning/stable/tutorials/installation/) for more details.
|
||||
|
||||
To start using Agent-lightning, check out our [documentation](https://microsoft.github.io/agent-lightning/) and [examples](./examples).
|
||||
|
||||
## ⚡ Articles
|
||||
|
||||
- 11/4/2025 [Tuning ANY AI agent with Tinker ✕ Agent-lightning](https://medium.com/@yugez/tuning-any-ai-agent-with-tinker-agent-lightning-part-1-1d8c9a397f0e) Medium. See also [Part 2](https://medium.com/@yugez/tuning-any-ai-agent-with-tinker-agent-lightning-part-2-332c5437f0dc).
|
||||
- 10/22/2025 [No More Retokenization Drift: Returning Token IDs via the OpenAI Compatible API Matters in Agent RL](https://blog.vllm.ai/2025/10/22/agent-lightning.html) vLLM blog. See also [Zhihu writeup](https://zhuanlan.zhihu.com/p/1965067274642785725).
|
||||
- 8/11/2025 [Training AI Agents to Write and Self-correct SQL with Reinforcement Learning](https://medium.com/@yugez/training-ai-agents-to-write-and-self-correct-sql-with-reinforcement-learning-571ed31281ad) Medium.
|
||||
- 8/5/2025 [Agent Lightning: Train ANY AI Agents with Reinforcement Learning](https://arxiv.org/abs/2508.03680) arXiv paper.
|
||||
@@ -67,10 +75,11 @@ No rewrites, no lock-in, just a clear path from first rollout to steady improvem
|
||||
| Workflow | Status |
|
||||
|----------|--------|
|
||||
| CPU Tests | [](https://github.com/microsoft/agent-lightning/actions/workflows/tests.yml) |
|
||||
| GPU Tests | [](https://github.com/microsoft/agent-lightning/actions/workflows/tests-full.yml) |
|
||||
| Full Tests | [](https://github.com/microsoft/agent-lightning/actions/workflows/badge-unit.yml) |
|
||||
| UI Tests | [](https://github.com/microsoft/agent-lightning/actions/workflows/dashboard.yml) |
|
||||
| Examples Integration | [](https://github.com/microsoft/agent-lightning/actions/workflows/badge-examples.yml) |
|
||||
| Latest Dependency Compatibility | [](https://github.com/microsoft/agent-lightning/actions/workflows/badge-latest.yml) |
|
||||
| Legacy Examples Compatibility | [](https://github.com/microsoft/agent-lightning/actions/workflows/examples-compat.yml) |
|
||||
| Legacy Examples Compatibility | [](https://github.com/microsoft/agent-lightning/actions/workflows/badge-compat.yml) |
|
||||
|
||||
## ⚡ Citation
|
||||
|
||||
@@ -90,7 +99,7 @@ If you find Agent Lightning useful in your research or projects, please cite our
|
||||
|
||||
## ⚡ Contributing
|
||||
|
||||
This project welcomes contributions and suggestions. Most contributions require you to agree to a Contributor License Agreement (CLA) declaring that you have the right to, and actually do, grant us the rights to use your contribution. For details, visit https://cla.opensource.microsoft.com.
|
||||
This project welcomes contributions and suggestions. Start by reading the [Contributing Guide](docs/community/contributing.md) for environment setup, branching conventions, and pull request expectations. Most contributions require you to agree to a Contributor License Agreement (CLA) declaring that you have the right to, and actually do, grant us the rights to use your contribution. For details, visit https://cla.opensource.microsoft.com.
|
||||
|
||||
When you submit a pull request, a CLA bot will automatically determine whether you need to provide a CLA and decorate the PR appropriately (e.g., status check, comment). Simply follow the instructions provided by the bot. You will only need to do this once across all repos using our CLA.
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
__version__ = "0.2.1"
|
||||
__version__ = "0.3.0"
|
||||
|
||||
from .adapter import *
|
||||
from .algorithm import *
|
||||
@@ -10,7 +10,9 @@ from .emitter import *
|
||||
from .execution import *
|
||||
from .litagent import *
|
||||
from .llm_proxy import *
|
||||
from .logging import *
|
||||
from .logging import configure_logger # deprecated # type: ignore
|
||||
from .logging import setup as setup_logging # type: ignore
|
||||
from .logging import setup_module as setup_module_logging # type: ignore
|
||||
from .runner import *
|
||||
from .server import AgentLightningServer # deprecated # type: ignore
|
||||
from .store import *
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from typing import Generic, List, TypeVar
|
||||
from typing import Generic, Sequence, TypeVar
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
@@ -66,7 +66,7 @@ class Adapter(Generic[T_from, T_to]):
|
||||
raise NotImplementedError("Adapter.adapt() is not implemented")
|
||||
|
||||
|
||||
class OtelTraceAdapter(Adapter[List[ReadableSpan], T_to], Generic[T_to]):
|
||||
class OtelTraceAdapter(Adapter[Sequence[ReadableSpan], T_to], Generic[T_to]):
|
||||
"""Base class for adapters that convert OpenTelemetry trace spans into other formats.
|
||||
|
||||
This specialization of [`Adapter`][agentlightning.Adapter] expects a list of
|
||||
@@ -84,7 +84,7 @@ class OtelTraceAdapter(Adapter[List[ReadableSpan], T_to], Generic[T_to]):
|
||||
"""
|
||||
|
||||
|
||||
class TraceAdapter(Adapter[List[Span], T_to], Generic[T_to]):
|
||||
class TraceAdapter(Adapter[Sequence[Span], T_to], Generic[T_to]):
|
||||
"""Base class for adapters that convert trace spans into other formats.
|
||||
|
||||
This class specializes [`Adapter`][agentlightning.Adapter] for working with
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING, Any, Dict, Generator, Iterable, List, Optional, TypedDict, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, Generator, Iterable, List, Optional, Sequence, TypedDict, Union, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
@@ -208,7 +208,7 @@ class TraceToMessages(TraceAdapter[List[OpenAIMessages]]):
|
||||
children of the associated completion span.
|
||||
"""
|
||||
|
||||
def get_tool_calls(self, completion: Span, all_spans: List[Span], /) -> Iterable[Dict[str, Any]]:
|
||||
def get_tool_calls(self, completion: Span, all_spans: Sequence[Span], /) -> Iterable[Dict[str, Any]]:
|
||||
"""Yield tool call payloads for a completion span.
|
||||
|
||||
Args:
|
||||
@@ -231,7 +231,7 @@ class TraceToMessages(TraceAdapter[List[OpenAIMessages]]):
|
||||
if tool_call:
|
||||
yield tool_call
|
||||
|
||||
def adapt(self, source: List[Span], /) -> List[OpenAIMessages]:
|
||||
def adapt(self, source: Sequence[Span], /) -> List[OpenAIMessages]:
|
||||
"""Transform trace spans into OpenAI chat payloads.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -6,7 +6,7 @@ import json
|
||||
import logging
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union, cast
|
||||
from typing import Any, Dict, List, Optional, Sequence, Tuple, Union, cast
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from pydantic import BaseModel
|
||||
@@ -429,7 +429,16 @@ class TraceTree:
|
||||
If we don't, when we want to select the LLM completion span with agent as filter.
|
||||
We will never get the correct span underneath.
|
||||
"""
|
||||
# If the current node has only one child, recursively repair its hierarchy directly.
|
||||
# This special-case handling is needed because when a trace is manually ended
|
||||
# (via agentops.end_trace), the AgentOps provider automatically wraps all spans
|
||||
# under an extra synthetic root node (e.g., "run_one.session").
|
||||
if len(self.children) == 1:
|
||||
self.children[0].repair_hierarchy()
|
||||
return
|
||||
|
||||
nodes_to_repair = list(self.children)
|
||||
|
||||
for repair_node in nodes_to_repair:
|
||||
if len(self.children) == 1:
|
||||
# If there is only one child, we don't need to repair the hierarchy.
|
||||
@@ -661,7 +670,7 @@ class TracerTraceToTriplet(TraceToTripletBase):
|
||||
trace_tree.visualize(filename, interested_span_match=interested_span_match)
|
||||
return trace_tree
|
||||
|
||||
def adapt(self, source: Union[List[Span], List[ReadableSpan]], /) -> List[Triplet]: # type: ignore
|
||||
def adapt(self, source: Union[Sequence[Span], Sequence[ReadableSpan]], /) -> List[Triplet]: # type: ignore
|
||||
"""Convert tracer spans into [`Triplet`][agentlightning.Triplet] trajectories.
|
||||
|
||||
Args:
|
||||
@@ -791,7 +800,7 @@ class LlmProxyTraceToTriplet(TraceToTripletBase):
|
||||
rid = attrs.get("gen_ai.response.id") or attrs.get("llm.hosted_vllm.id")
|
||||
return str(rid) if isinstance(rid, str) and rid else None
|
||||
|
||||
def adapt(self, source: List[Span], /) -> List[Triplet]: # type: ignore
|
||||
def adapt(self, source: Sequence[Span], /) -> List[Triplet]: # type: ignore
|
||||
"""Convert LLM Proxy spans into [`Triplet`][agentlightning.Triplet] trajectories.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -20,6 +20,7 @@ from openai import AsyncOpenAI
|
||||
|
||||
from agentlightning.adapter.messages import TraceToMessages
|
||||
from agentlightning.algorithm.base import Algorithm
|
||||
from agentlightning.algorithm.utils import batch_iter_over_dataset
|
||||
from agentlightning.reward import find_final_reward
|
||||
from agentlightning.types import Dataset, NamedResources, PromptTemplate, Rollout, RolloutMode, RolloutStatus
|
||||
|
||||
@@ -56,41 +57,6 @@ APPLY_EDIT_PROMPT_FILES = [
|
||||
]
|
||||
|
||||
|
||||
def batch_iter_over_dataset(dataset: Dataset[T_task], batch_size: int) -> Iterator[Sequence[T_task]]:
|
||||
"""
|
||||
Create an infinite iterator that yields batches from the dataset.
|
||||
|
||||
When batch_size >= dataset size, yields the entire shuffled dataset repeatedly.
|
||||
When batch_size < dataset size, yields batches of the specified size, reshuffling
|
||||
after each complete pass through the dataset.
|
||||
|
||||
Args:
|
||||
dataset: The dataset to iterate over.
|
||||
batch_size: The desired batch size.
|
||||
|
||||
Yields:
|
||||
Sequences of tasks from the dataset. Each task appears at most once per epoch.
|
||||
"""
|
||||
if batch_size >= len(dataset):
|
||||
while True:
|
||||
dataset_copy = [dataset[i] for i in range(len(dataset))]
|
||||
random.shuffle(dataset_copy)
|
||||
yield dataset_copy
|
||||
|
||||
else:
|
||||
current_batch: List[int] = []
|
||||
while True:
|
||||
indices = list(range(len(dataset)))
|
||||
random.shuffle(indices)
|
||||
for index in indices:
|
||||
if index in current_batch:
|
||||
continue
|
||||
current_batch.append(index)
|
||||
if len(current_batch) == batch_size:
|
||||
yield [dataset[index] for index in current_batch]
|
||||
current_batch = []
|
||||
|
||||
|
||||
class APO(Algorithm, Generic[T_task]):
|
||||
"""Automatic Prompt Optimization (APO) algorithm using textual gradients and beam search.
|
||||
|
||||
|
||||
@@ -143,7 +143,7 @@ class Baseline(FastAlgorithm):
|
||||
store = self.get_store()
|
||||
|
||||
for index in train_indices + val_indices:
|
||||
queuing_rollouts = await store.query_rollouts(status=["queuing", "requeuing"])
|
||||
queuing_rollouts = await store.query_rollouts(status_in=["queuing", "requeuing"])
|
||||
if len(queuing_rollouts) <= 1:
|
||||
# Only enqueue a new rollout when there is at most 1 rollout in the queue.
|
||||
sample = dataset[index]
|
||||
@@ -222,7 +222,7 @@ class Baseline(FastAlgorithm):
|
||||
f"Processing index {index}. {len(train_indices)} train indices and {len(val_indices)} val indices in total."
|
||||
)
|
||||
while True:
|
||||
queuing_rollouts = await store.query_rollouts(status=["queuing", "requeuing"])
|
||||
queuing_rollouts = await store.query_rollouts(status_in=["queuing", "requeuing"])
|
||||
if len(queuing_rollouts) <= self.max_queue_length:
|
||||
# Only enqueue a new rollout when there is at most "max_queue_length" rollout in the queue.
|
||||
sample = concatenated_dataset[index]
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import random
|
||||
from typing import Iterator, List, Sequence, TypeVar
|
||||
|
||||
from agentlightning.types import Dataset
|
||||
|
||||
T_task = TypeVar("T_task")
|
||||
|
||||
|
||||
def batch_iter_over_dataset(dataset: Dataset[T_task], batch_size: int) -> Iterator[Sequence[T_task]]:
|
||||
"""
|
||||
Create an infinite iterator that yields batches from the dataset.
|
||||
|
||||
When batch_size >= dataset size, yields the entire shuffled dataset repeatedly.
|
||||
When batch_size < dataset size, yields batches of the specified size, reshuffling
|
||||
after each complete pass through the dataset.
|
||||
|
||||
Args:
|
||||
dataset: The dataset to iterate over.
|
||||
batch_size: The desired batch size.
|
||||
|
||||
Yields:
|
||||
Sequences of tasks from the dataset. Each task appears at most once per epoch.
|
||||
"""
|
||||
if batch_size >= len(dataset):
|
||||
while True:
|
||||
dataset_copy = [dataset[i] for i in range(len(dataset))]
|
||||
random.shuffle(dataset_copy)
|
||||
yield dataset_copy
|
||||
|
||||
else:
|
||||
current_batch: List[int] = []
|
||||
while True:
|
||||
indices = list(range(len(dataset)))
|
||||
random.shuffle(indices)
|
||||
for index in indices:
|
||||
if index in current_batch:
|
||||
continue
|
||||
current_batch.append(index)
|
||||
if len(current_batch) == batch_size:
|
||||
yield [dataset[index] for index in current_batch]
|
||||
current_batch = []
|
||||
@@ -99,6 +99,8 @@ class VERL(Algorithm):
|
||||
|
||||
# Merge your dict overrides
|
||||
override_conf = OmegaConf.create(config)
|
||||
# Allow adding new fields
|
||||
OmegaConf.set_struct(base_cfg, False)
|
||||
self.config = OmegaConf.merge(base_cfg, override_conf)
|
||||
|
||||
def run(
|
||||
|
||||
@@ -1,30 +0,0 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import time
|
||||
from typing import Iterable
|
||||
|
||||
from agentlightning.instrumentation.agentops import AgentOpsServerManager
|
||||
|
||||
|
||||
def main(argv: Iterable[str] | None = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="Start AgentOps server")
|
||||
parser.add_argument("--daemon", action="store_true", help="Run server as a daemon")
|
||||
parser.add_argument("--port", type=int, default=8002, help="Port to run the server on")
|
||||
args = parser.parse_args(list(argv) if argv is not None else None)
|
||||
|
||||
manager = AgentOpsServerManager(daemon=args.daemon, port=args.port)
|
||||
try:
|
||||
manager.start()
|
||||
# Wait forever
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
manager.stop()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -9,7 +9,7 @@ import asyncio
|
||||
import logging
|
||||
from typing import Iterable
|
||||
|
||||
from agentlightning.logging import configure_logger
|
||||
from agentlightning import setup_logging
|
||||
from agentlightning.store.client_server import LightningStoreServer
|
||||
from agentlightning.store.memory import InMemoryLightningStore
|
||||
|
||||
@@ -19,12 +19,30 @@ logger = logging.getLogger(__name__)
|
||||
def main(argv: Iterable[str] | None = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="Run a LightningStore server")
|
||||
parser.add_argument("--port", type=int, default=4747, help="Port to run the server on")
|
||||
parser.add_argument(
|
||||
"--cors-origin",
|
||||
dest="cors_origins",
|
||||
action="append",
|
||||
help="Allowed CORS origin. Repeat for multiple origins. Use '*' to allow all origins.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
default="INFO",
|
||||
choices=["DEBUG", "INFO", "WARNING", "ERROR"],
|
||||
help="Configure the logging level for the store.",
|
||||
)
|
||||
args = parser.parse_args(list(argv) if argv is not None else None)
|
||||
|
||||
configure_logger()
|
||||
setup_logging(args.log_level)
|
||||
|
||||
store = InMemoryLightningStore()
|
||||
server = LightningStoreServer(store, host="0.0.0.0", port=args.port)
|
||||
server = LightningStoreServer(
|
||||
store,
|
||||
host="0.0.0.0",
|
||||
port=args.port,
|
||||
cors_allow_origins=args.cors_origins,
|
||||
launch_mode="asyncio",
|
||||
)
|
||||
try:
|
||||
asyncio.run(server.run_forever())
|
||||
except RuntimeError as exc:
|
||||
|
||||
@@ -332,7 +332,9 @@ class DevTaskLoader(AgentLightningClient):
|
||||
if isinstance(resources, ResourcesUpdate):
|
||||
self._resources_update = resources
|
||||
else:
|
||||
self._resources_update = ResourcesUpdate(resources_id="local", resources=resources)
|
||||
self._resources_update = ResourcesUpdate(
|
||||
resources_id="local", resources=resources, create_time=time.time(), update_time=time.time(), version=1
|
||||
)
|
||||
|
||||
# Store rollouts posted back to the loader for easy debugging of local runs
|
||||
self._rollouts: List[RolloutLegacy] = []
|
||||
|
||||
@@ -129,12 +129,13 @@ def reward(fn: FnType) -> FnType:
|
||||
return wrapper # type: ignore
|
||||
|
||||
|
||||
def emit_reward(reward: float) -> ReadableSpan:
|
||||
def emit_reward(reward: float, auto_export: bool = True) -> ReadableSpan:
|
||||
"""Emit a reward value as an OpenTelemetry span.
|
||||
|
||||
Args:
|
||||
reward: Numeric reward to record. Integers and booleans are converted to
|
||||
floating point numbers for consistency.
|
||||
auto_export: Whether to export the span automatically.
|
||||
|
||||
Returns:
|
||||
Readable span capturing the recorded reward.
|
||||
@@ -150,7 +151,7 @@ def emit_reward(reward: float) -> ReadableSpan:
|
||||
raise ValueError(f"Reward must be a number, got: {type(reward)}")
|
||||
|
||||
# TODO: This should use the tracer from current context by tracer
|
||||
tracer = get_tracer()
|
||||
tracer = get_tracer(use_active_span_processor=auto_export)
|
||||
span = tracer.start_span(SpanNames.REWARD.value, attributes={"reward": reward})
|
||||
# Do nothing; it's just a number
|
||||
with span:
|
||||
|
||||
@@ -2,13 +2,22 @@
|
||||
|
||||
"""Utilities shared across emitter implementations."""
|
||||
|
||||
from typing import cast
|
||||
from warnings import filterwarnings
|
||||
|
||||
import opentelemetry.trace as trace_api
|
||||
from opentelemetry.sdk.trace import SpanLimits, SynchronousMultiSpanProcessor, Tracer
|
||||
from opentelemetry.sdk.trace import TracerProvider as TracerProviderImpl
|
||||
from opentelemetry.sdk.util.instrumentation import InstrumentationInfo, InstrumentationScope
|
||||
from opentelemetry.trace import get_tracer_provider
|
||||
|
||||
|
||||
def get_tracer() -> trace_api.Tracer:
|
||||
def get_tracer(use_active_span_processor: bool = True) -> trace_api.Tracer:
|
||||
"""Resolve the OpenTelemetry tracer configured for Agent Lightning.
|
||||
|
||||
Args:
|
||||
use_active_span_processor: Whether to use the active span processor.
|
||||
|
||||
Returns:
|
||||
OpenTelemetry tracer tagged with the `agentlightning` instrumentation name.
|
||||
|
||||
@@ -18,5 +27,31 @@ def get_tracer() -> trace_api.Tracer:
|
||||
if hasattr(trace_api, "_TRACER_PROVIDER") and trace_api._TRACER_PROVIDER is None: # type: ignore[attr-defined]
|
||||
raise RuntimeError("Tracer is not initialized. Cannot emit a meaningful span.")
|
||||
|
||||
tracer_provider = get_tracer_provider()
|
||||
return tracer_provider.get_tracer("agentlightning")
|
||||
tracer_provider = cast(TracerProviderImpl, get_tracer_provider())
|
||||
|
||||
if use_active_span_processor:
|
||||
return tracer_provider.get_tracer("agentlightning")
|
||||
|
||||
else:
|
||||
filterwarnings(
|
||||
"ignore",
|
||||
message=r"You should use InstrumentationScope. Deprecated since version 1.11.1.",
|
||||
category=DeprecationWarning,
|
||||
module="opentelemetry.sdk.trace",
|
||||
)
|
||||
|
||||
return Tracer(
|
||||
tracer_provider.sampler,
|
||||
tracer_provider.resource,
|
||||
# We use an empty span processor to avoid emitting spans to the tracer
|
||||
SynchronousMultiSpanProcessor(),
|
||||
tracer_provider.id_generator,
|
||||
InstrumentationInfo("agentlightning", "", ""), # type: ignore
|
||||
SpanLimits(),
|
||||
InstrumentationScope(
|
||||
"agentlightning",
|
||||
"",
|
||||
"",
|
||||
{},
|
||||
),
|
||||
)
|
||||
|
||||
@@ -67,8 +67,8 @@ class ClientServerExecutionStrategy(ExecutionStrategy):
|
||||
server_host: str | None = None,
|
||||
server_port: int | None = None,
|
||||
n_runners: int = 1,
|
||||
graceful_timeout: float = 5.0,
|
||||
terminate_timeout: float = 5.0,
|
||||
graceful_timeout: float = 10.0,
|
||||
terminate_timeout: float = 10.0,
|
||||
main_process: Literal["algorithm", "runner"] = "algorithm",
|
||||
managed_store: bool | None = None,
|
||||
) -> None:
|
||||
|
||||
@@ -4,28 +4,76 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import multiprocessing
|
||||
import signal
|
||||
import socket
|
||||
import time
|
||||
from typing import Any, Callable, no_type_check
|
||||
|
||||
import flask
|
||||
import requests
|
||||
import setproctitle
|
||||
from agentops.client.api import V3Client, V4Client
|
||||
from agentops.client.api.types import AuthTokenResponse
|
||||
from agentops.sdk.exporters import AuthenticatedOTLPExporter
|
||||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import OTLPMetricExporter
|
||||
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
||||
from opentelemetry.sdk.metrics.export import MetricExportResult
|
||||
|
||||
from agentlightning.utils.otlp import LightningStoreOTLPExporter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
__all__ = [
|
||||
"instrument_agentops",
|
||||
"uninstrument_agentops",
|
||||
"agentops_local_server",
|
||||
"AgentOpsServerManager",
|
||||
]
|
||||
|
||||
# Module-level storage for originals
|
||||
_original_handle_chat_attributes: Callable[..., Any] | None = None
|
||||
_original_handle_response: Callable[..., Any] | None = None
|
||||
_agentops_service_enabled = False
|
||||
|
||||
|
||||
def enable_agentops_service(enabled: bool = True) -> None:
|
||||
"""
|
||||
Enable or disable communication with the AgentOps service.
|
||||
|
||||
By default, AgentOps exporters and clients will run in local mode
|
||||
and will NOT attempt to communicate with the remote AgentOps service.
|
||||
|
||||
Args:
|
||||
enabled: If True, enable all AgentOps exporters and clients.
|
||||
All exporters and clients will operate in normal mode and send data
|
||||
to the [AgentOps service](https://www.agentops.ai).
|
||||
"""
|
||||
global _agentops_service_enabled
|
||||
_agentops_service_enabled = enabled
|
||||
logger.info(f"AgentOps service enabled is set to {enabled}.")
|
||||
|
||||
|
||||
def _patch_exporters():
|
||||
import agentops.client.api
|
||||
import agentops.sdk.core
|
||||
|
||||
agentops.sdk.core.AuthenticatedOTLPExporter = BypassableAuthenticatedOTLPExporter # type: ignore
|
||||
agentops.sdk.core.OTLPMetricExporter = BypassableOTLPMetricExporter
|
||||
if hasattr(agentops.sdk.core, "OTLPSpanExporter"):
|
||||
agentops.sdk.core.OTLPSpanExporter = BypassableOTLPSpanExporter # type: ignore
|
||||
agentops.client.api.V3Client = BypassableV3Client
|
||||
agentops.client.api.V4Client = BypassableV4Client
|
||||
|
||||
|
||||
def _unpatch_exporters():
|
||||
import agentops.client.api
|
||||
import agentops.sdk.core
|
||||
|
||||
agentops.sdk.core.AuthenticatedOTLPExporter = AuthenticatedOTLPExporter # type: ignore
|
||||
agentops.sdk.core.OTLPMetricExporter = OTLPMetricExporter
|
||||
if hasattr(agentops.sdk.core, "OTLPSpanExporter"):
|
||||
agentops.sdk.core.OTLPSpanExporter = OTLPSpanExporter # type: ignore
|
||||
agentops.client.api.V3Client = V3Client
|
||||
agentops.client.api.V4Client = V4Client
|
||||
|
||||
|
||||
def _unwrap_legacy_response(response: Any) -> Any:
|
||||
if hasattr(response, "parse") and callable(response.parse):
|
||||
return response.parse()
|
||||
return response
|
||||
|
||||
|
||||
def _patch_new_agentops():
|
||||
@@ -44,6 +92,11 @@ def _patch_new_agentops():
|
||||
@no_type_check
|
||||
def _handle_chat_attributes_with_tokens(args=None, kwargs=None, return_value=None, **kws): # type: ignore
|
||||
attributes = _original_handle_chat_attributes(args=args, kwargs=kwargs, return_value=return_value, **kws)
|
||||
|
||||
# In some cases, response is a openai._legacy_response.LegacyAPIResponse (e.g., LiteLLM, or LangChain),
|
||||
# This is created by client.with_raw_response.create()
|
||||
return_value = _unwrap_legacy_response(return_value)
|
||||
|
||||
if (
|
||||
return_value is not None
|
||||
and hasattr(return_value, "prompt_token_ids")
|
||||
@@ -89,20 +142,6 @@ def _patch_new_agentops():
|
||||
[logprob.model_dump() for logprob in first_choice.logprobs.refusal]
|
||||
)
|
||||
|
||||
# For LiteLLM, response is a openai._legacy_response.LegacyAPIResponse
|
||||
if (
|
||||
return_value is not None
|
||||
and hasattr(return_value, "http_response")
|
||||
and return_value.http_response is not None
|
||||
and hasattr(return_value.http_response, "json")
|
||||
):
|
||||
json_data = return_value.http_response.json()
|
||||
if isinstance(json_data, dict):
|
||||
if json_data.get("prompt_token_ids") is not None:
|
||||
attributes["prompt_token_ids"] = list(json_data["prompt_token_ids"])
|
||||
if json_data.get("response_token_ids") is not None:
|
||||
attributes["response_token_ids"] = list(json_data["response_token_ids"][0])
|
||||
|
||||
return attributes
|
||||
|
||||
agentops.instrumentation.providers.openai.wrappers.chat.handle_chat_attributes = _handle_chat_attributes_with_tokens
|
||||
@@ -173,6 +212,8 @@ def instrument_agentops():
|
||||
Instrument agentops to capture token IDs.
|
||||
Automatically detects and uses the appropriate patching method based on the installed agentops version.
|
||||
"""
|
||||
_patch_exporters()
|
||||
|
||||
# Try newest version first (tested for 0.4.16)
|
||||
try:
|
||||
return _patch_new_agentops()
|
||||
@@ -192,6 +233,8 @@ def instrument_agentops():
|
||||
|
||||
def uninstrument_agentops():
|
||||
"""Uninstrument agentops to stop capturing token IDs."""
|
||||
_unpatch_exporters()
|
||||
|
||||
try:
|
||||
_unpatch_new_agentops()
|
||||
except Exception:
|
||||
@@ -202,114 +245,70 @@ def uninstrument_agentops():
|
||||
pass
|
||||
|
||||
|
||||
def agentops_local_server():
|
||||
class BypassableAuthenticatedOTLPExporter(LightningStoreOTLPExporter, AuthenticatedOTLPExporter):
|
||||
"""
|
||||
Returns a Flask app that can be used to test agentops integration.
|
||||
This server provides endpoints for token fetching and a catch-all endpoint.
|
||||
AuthenticatedOTLPExporter with switchable service control.
|
||||
|
||||
When `_agentops_service_enabled` is False, skip export and return success.
|
||||
"""
|
||||
app = flask.Flask(__name__)
|
||||
|
||||
@app.route("/v3/auth/token", methods=["POST"])
|
||||
def fetch_token(): # type: ignore
|
||||
return {"token": "dummy", "project_id": "dummy"}
|
||||
|
||||
@app.route("/", defaults={"path": ""}, methods=["GET", "POST"])
|
||||
@app.route("/<path:path>", methods=["GET", "POST"])
|
||||
def catch_all(path: str): # type: ignore
|
||||
return {"path": path}
|
||||
|
||||
return app
|
||||
def should_bypass(self) -> bool:
|
||||
return not _agentops_service_enabled
|
||||
|
||||
|
||||
def _run_server(**kwargs: Any): # type: ignore
|
||||
class BypassableOTLPMetricExporter(OTLPMetricExporter):
|
||||
"""
|
||||
Internal function to run the Flask server.
|
||||
This is used to avoid issues with multiprocessing and Flask's reloader.
|
||||
OTLPMetricExporter with switchable service control.
|
||||
When `_agentops_service_enabled` is False, skip export and return success.
|
||||
"""
|
||||
signal.signal(signal.SIGINT, signal.SIG_IGN) # Ignore SIGINT in worker processes
|
||||
setproctitle.setproctitle(multiprocessing.current_process().name)
|
||||
app = agentops_local_server()
|
||||
app.run(**kwargs)
|
||||
|
||||
|
||||
class AgentOpsServerManager:
|
||||
"""Manages a AgentOps local server to bypass the online service of AgentOps."""
|
||||
|
||||
def __init__(self, daemon: bool = True, port: int | None = None):
|
||||
self.server_process: multiprocessing.Process | None = None
|
||||
self.server_port = port
|
||||
self.daemon = daemon
|
||||
logger.info("AgentOpsServerManager initialized.")
|
||||
|
||||
def _find_available_port(self) -> int:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
def start(self):
|
||||
if self.server_process and self.server_process.is_alive():
|
||||
logger.warning("AgentOps server process appears to be already running.")
|
||||
return
|
||||
|
||||
if self.server_port is None:
|
||||
self.server_port = self._find_available_port()
|
||||
|
||||
logger.info(f"Starting AgentOps local server on port {self.server_port}...")
|
||||
|
||||
self.server_process = multiprocessing.Process(
|
||||
target=_run_server,
|
||||
kwargs={"host": "127.0.0.1", "port": self.server_port, "use_reloader": False, "debug": False},
|
||||
daemon=self.daemon,
|
||||
name="AgentLightning-AgentOpsServer",
|
||||
)
|
||||
self.server_process.start()
|
||||
logger.info(
|
||||
f"AgentOps local server process (PID: {self.server_process.pid}) started, targeting port {self.server_port}."
|
||||
)
|
||||
for attempt in range(20): # 10 seconds total
|
||||
time.sleep(0.5) # Brief wait for server to start up
|
||||
try:
|
||||
result = requests.get(f"http://127.0.0.1:{self.server_port}/")
|
||||
if result.status_code == 200:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.debug(f"Error checking AgentOps server: {e}")
|
||||
logger.warning(f"AgentOps still not ready after {attempt} attempts. Retrying...")
|
||||
def export(self, *args: Any, **kwargs: Any) -> MetricExportResult:
|
||||
if _agentops_service_enabled:
|
||||
return super().export(*args, **kwargs) # type: ignore[reportUnknownMemberType]
|
||||
else:
|
||||
logger.error(f"AgentOps local server failed to start or exited prematurely.")
|
||||
return
|
||||
logger.debug("SwitchableOTLPMetricExporter is switched off, skipping export.")
|
||||
return MetricExportResult.SUCCESS
|
||||
|
||||
if not self.server_process.is_alive():
|
||||
logger.error(f"AgentOps local server failed to start or exited prematurely.")
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
if self.server_process and self.server_process.is_alive():
|
||||
return True
|
||||
return False
|
||||
class BypassableOTLPSpanExporter(LightningStoreOTLPExporter):
|
||||
"""
|
||||
OTLPSpanExporter with switchable service control.
|
||||
When `_agentops_service_enabled` is False, skip export and return success.
|
||||
|
||||
def stop(self):
|
||||
if self.server_process is not None and self.server_process.is_alive():
|
||||
logger.info(f"Stopping AgentOps local server (PID: {self.server_process.pid})...")
|
||||
self.server_process.terminate() # Send SIGTERM
|
||||
self.server_process.join(timeout=5) # Wait for clean exit
|
||||
if self.server_process.is_alive():
|
||||
logger.warning(
|
||||
f"AgentOps server (PID: {self.server_process.pid}) did not terminate gracefully, killing..."
|
||||
)
|
||||
self.server_process.kill() # Force kill
|
||||
self.server_process.join(timeout=10) # Wait for kill
|
||||
self.server_process = None
|
||||
logger.info(f"AgentOps local server stopped.")
|
||||
This is used instead of BypassableAuthenticatedOTLPExporter on legacy AgentOps versions.
|
||||
"""
|
||||
|
||||
def should_bypass(self) -> bool:
|
||||
return not _agentops_service_enabled
|
||||
|
||||
|
||||
class BypassableV3Client(V3Client):
|
||||
"""
|
||||
V3Client with toggleable authentication calls.
|
||||
Returns dummy auth response when `_agentops_service_enabled` is False.
|
||||
"""
|
||||
|
||||
# Temporary synchronous override of fetch_auth_token for mock purposes.
|
||||
def fetch_auth_token(self, *args: Any, **kwargs: Any) -> AuthTokenResponse: # type: ignore[override]
|
||||
if _agentops_service_enabled:
|
||||
return super().fetch_auth_token(*args, **kwargs) # type: ignore[override]
|
||||
else:
|
||||
logger.info("AgentOps local server was not running or already stopped.")
|
||||
logger.debug("SwitchableV3Client is switched off, skipping fetch_auth_token request.")
|
||||
return AuthTokenResponse(token="dummy", project_id="dummy")
|
||||
|
||||
def get_port(self) -> int | None:
|
||||
# Check liveness again in case it died since start()
|
||||
if self.is_alive() and self.server_port is not None:
|
||||
return self.server_port
|
||||
# If called after server stopped or failed, port might be stale or None
|
||||
if self.server_port is not None and (self.server_process is None or not self.server_process.is_alive()):
|
||||
logger.warning(
|
||||
f"AgentOps server port {self.server_port} is stored, but server process is not alive. Returning stored port."
|
||||
)
|
||||
return self.server_port
|
||||
|
||||
class BypassableV4Client(V4Client):
|
||||
"""
|
||||
V4Client with toggleable post requests.
|
||||
Returns dummy response when `_agentops_service_enabled` is False.
|
||||
"""
|
||||
|
||||
def post(self, *args: Any, **kwargs: Any) -> requests.Response:
|
||||
if _agentops_service_enabled:
|
||||
return super().post(*args, **kwargs)
|
||||
else:
|
||||
logger.debug("SwitchableV4Client is switched off, skipping post request.")
|
||||
response = requests.Response()
|
||||
response.status_code = 200
|
||||
response._content = b"{}"
|
||||
return response
|
||||
|
||||
+703
-194
File diff suppressed because it is too large
Load Diff
+329
-13
@@ -1,10 +1,18 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import sys
|
||||
import warnings
|
||||
from logging.config import dictConfig
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
__all__ = ["configure_logger"]
|
||||
from rich.console import Console
|
||||
|
||||
__all__ = ["setup", "configure_logger", "setup_module"]
|
||||
|
||||
|
||||
def configure_logger(level: int = logging.INFO, name: str = "agentlightning") -> logging.Logger:
|
||||
@@ -15,6 +23,10 @@ def configure_logger(level: int = logging.INFO, name: str = "agentlightning") ->
|
||||
not propagate to the root logger, preventing duplicate log emission when
|
||||
applications compose multiple logging configurations.
|
||||
|
||||
!!! danger
|
||||
|
||||
This function is deprecated in favor of [`setup_logging`][agentlightning.setup_logging].
|
||||
|
||||
Args:
|
||||
level: Logging level applied both to the logger and the installed
|
||||
handler. Defaults to `logging.INFO`.
|
||||
@@ -32,23 +44,327 @@ def configure_logger(level: int = logging.INFO, name: str = "agentlightning") ->
|
||||
logger.info("agent-lightning is ready!")
|
||||
```
|
||||
"""
|
||||
warnings.warn("This function is deprecated in favor of `setup_logging`.", DeprecationWarning, stacklevel=2)
|
||||
|
||||
return setup_module(level=level, name=name, console=True, color=True, propagate=False)
|
||||
|
||||
|
||||
DEFAULT_FORMAT = "%(asctime)s [%(levelname)s] (Process-%(process)d %(name)s) %(message)s"
|
||||
DATE_FORMAT = "%H:%M:%S"
|
||||
|
||||
|
||||
def _to_level_value(lvl: int | str) -> int:
|
||||
if isinstance(lvl, int):
|
||||
return lvl
|
||||
val = getattr(logging, str(lvl).upper(), None)
|
||||
if val is None:
|
||||
raise ValueError(f"Invalid log level: {lvl}")
|
||||
return val
|
||||
|
||||
|
||||
def _ensure_file_handler(
|
||||
logger: logging.Logger,
|
||||
filename: str,
|
||||
*,
|
||||
level: int,
|
||||
formatter: Optional[logging.Formatter],
|
||||
) -> None:
|
||||
"""Attach a FileHandler to `logger` for `filename` if it doesn't already exist."""
|
||||
abspath = os.path.abspath(filename)
|
||||
|
||||
# Avoid duplicates
|
||||
for h in logger.handlers:
|
||||
if isinstance(h, logging.FileHandler) and getattr(h, "baseFilename", None) == abspath:
|
||||
return
|
||||
|
||||
# Ensure directory exists
|
||||
dirname = os.path.dirname(abspath)
|
||||
if dirname:
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
|
||||
fh = logging.FileHandler(abspath, encoding="utf-8")
|
||||
fh.setLevel(level)
|
||||
if formatter is not None:
|
||||
fh.setFormatter(formatter)
|
||||
else:
|
||||
fh.setFormatter(logging.Formatter(DEFAULT_FORMAT, DATE_FORMAT))
|
||||
|
||||
logger.addHandler(fh)
|
||||
|
||||
|
||||
def setup(
|
||||
level: int | str = "INFO",
|
||||
*,
|
||||
console: bool = True,
|
||||
color: bool | Dict[str, Any] = True,
|
||||
propagate: bool = False,
|
||||
disable_existing_loggers: bool = False,
|
||||
capture_warnings: bool = False,
|
||||
submodule_levels: Optional[dict[str, int | str]] = None,
|
||||
extra_handlers: Optional[list[logging.Handler]] = None,
|
||||
formatter: Optional[logging.Formatter] = None,
|
||||
apply_to: Optional[list[str]] = None,
|
||||
files: Optional[str | dict[str, str]] = None,
|
||||
) -> None:
|
||||
"""Configures logging for the `agentlightning` logger hierarchy.
|
||||
|
||||
This function provides a one-stop setup utility for configuring the
|
||||
`agentlightning` root logger and optionally its submodules or external
|
||||
loggers. It supports console logging, colored rich output, per-submodule
|
||||
log levels, and optional handler/formatter injection.
|
||||
|
||||
The setup is intentionally isolated: it does not modify the global root
|
||||
logger or loggers belonging to other libraries unless explicitly directed
|
||||
via `apply_to`.
|
||||
|
||||
Args:
|
||||
level:
|
||||
Logging level for the base `agentlightning` logger. Accepts either
|
||||
an integer (e.g., `logging.DEBUG`) or a string level name
|
||||
(e.g., `"INFO"`). Defaults to `"INFO"`.
|
||||
console:
|
||||
Whether to attach a console handler to the logger. Defaults to
|
||||
`True`.
|
||||
color:
|
||||
Enables rich-formatted output using `RichHandler` when `True`
|
||||
or a configuration dict. If `False`, a plain text formatter is
|
||||
used instead. Defaults to `True`.
|
||||
propagate:
|
||||
Whether `agentlightning` logs should propagate to ancestor
|
||||
loggers. Defaults to `False`.
|
||||
disable_existing_loggers:
|
||||
Passed to `logging.config.dictConfig`. If `True`, disables all
|
||||
existing configured loggers before applying this configuration.
|
||||
Defaults to `False`.
|
||||
capture_warnings:
|
||||
If `True`, redirects Python `warnings` emitted via the `warnings`
|
||||
module into the logging system. Defaults to `False`.
|
||||
submodule_levels:
|
||||
Mapping of submodule logger names to logging levels. If a specified
|
||||
submodule level is more verbose than the base level, a warning is emitted.
|
||||
extra_handlers:
|
||||
A list of user-provided handlers to attach to the `agentlightning` logger.
|
||||
Handlers are added idempotently; duplicates are not reattached.
|
||||
formatter:
|
||||
A formatter to apply to any handler under `agentlightning` that does not
|
||||
already have one assigned. Useful for customizing output without overwriting
|
||||
formatters on custom handlers.
|
||||
apply_to:
|
||||
A list of additional logger names to configure identically to
|
||||
`agentlightning` base logger. Their handlers are replaced with copies of the base
|
||||
handlers, and propagation is disabled to avoid duplicate log emission.
|
||||
files:
|
||||
If a string, attach a FileHandler to the base `agentlightning` logger.
|
||||
If a dict, for each `(logger_name, filename)` pair, attach a FileHandler
|
||||
directly to that logger.
|
||||
Each file handler should use the logger's effective level at creation.
|
||||
|
||||
Notes:
|
||||
* On Windows, this function forces UTF-8 mode in the console to prevent
|
||||
issues with rich output or special characters.
|
||||
* Submodule loggers can generate records below the handler's emission
|
||||
threshold. Whether such records appear depends on both the logger's
|
||||
level and the handler's level.
|
||||
* `apply_to` loggers inherit the same handlers but do not propagate
|
||||
upward, yielding isolated, consistent behavior.
|
||||
|
||||
Examples:
|
||||
Basic setup:
|
||||
|
||||
>>> setup()
|
||||
|
||||
Enabling debug mode with no color:
|
||||
|
||||
>>> setup(level="DEBUG", color=False)
|
||||
|
||||
Overriding specific submodule levels:
|
||||
|
||||
>>> setup(submodule_levels={"agentlightning.io": "DEBUG"})
|
||||
|
||||
Attaching an additional file handler:
|
||||
|
||||
>>> fh = logging.FileHandler("app.log")
|
||||
>>> setup(extra_handlers=[fh])
|
||||
"""
|
||||
# Ensure UTF-8 encoding on Windows consoles
|
||||
# Note: This change does not fully represent support for execution under the windown system.
|
||||
# Note: This change does not fully represent support for execution under the windows system.
|
||||
# It only fixes console printing issues caused by special characters.
|
||||
# TODO: More comprehensive Windows support may be needed in the future.
|
||||
if platform.system() == "Windows":
|
||||
os.environ["PYTHONUTF8"] = "1"
|
||||
|
||||
logger = logging.getLogger(name)
|
||||
logger.handlers.clear() # clear existing handlers
|
||||
base_logger = setup_module(
|
||||
level,
|
||||
name="agentlightning",
|
||||
console=console,
|
||||
color=color,
|
||||
propagate=propagate,
|
||||
disable_existing_loggers=disable_existing_loggers,
|
||||
)
|
||||
|
||||
# log to stdout
|
||||
handler = logging.StreamHandler()
|
||||
handler.setLevel(level)
|
||||
formatter = logging.Formatter("%(asctime)s [%(levelname)s] (Process-%(process)d %(name)s) %(message)s")
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(level)
|
||||
logger.propagate = False # prevent double logging
|
||||
return logger
|
||||
base_level_value = base_logger.level
|
||||
|
||||
# Apply user-provided formatter (only to handlers without one,
|
||||
# so we don't clobber custom extra_handlers)
|
||||
if formatter is not None:
|
||||
for h in base_logger.handlers:
|
||||
if h.formatter is None:
|
||||
h.setFormatter(formatter)
|
||||
|
||||
# Attach user-provided handler(s) if any, idempotently
|
||||
if extra_handlers:
|
||||
for h in extra_handlers:
|
||||
if h not in base_logger.handlers:
|
||||
base_logger.addHandler(h)
|
||||
|
||||
# Per-submodule levels
|
||||
if submodule_levels:
|
||||
for name, lvl in submodule_levels.items():
|
||||
sub_level = _to_level_value(lvl)
|
||||
|
||||
# Emit a warning if submodule level is lower (more verbose) than the global/base level
|
||||
if sub_level < base_level_value:
|
||||
base_logger.warning(
|
||||
"Submodule logger '%s' level %s (%s) is more verbose than base "
|
||||
"logger level %s (%s). Records below the base level may still be "
|
||||
"filtered out by handlers depending on their own levels.",
|
||||
name,
|
||||
lvl,
|
||||
sub_level,
|
||||
logging.getLevelName(base_level_value),
|
||||
base_level_value,
|
||||
)
|
||||
|
||||
# The logger will *create* records down to the logger's level, but a handler
|
||||
# with a higher level will still drop anything below its own threshold.
|
||||
# Effective emission is gated by both: record.level >= logger.level AND handler.level.
|
||||
logging.getLogger(name).setLevel(lvl)
|
||||
|
||||
# Attach file handlers if requested
|
||||
if files is not None:
|
||||
if isinstance(files, str):
|
||||
# Single file for the entire `agentlightning` hierarchy.
|
||||
_ensure_file_handler(
|
||||
logger=base_logger,
|
||||
filename=files,
|
||||
level=base_level_value,
|
||||
formatter=formatter,
|
||||
)
|
||||
else:
|
||||
# Per-logger files
|
||||
for logger_name, filename in files.items():
|
||||
lg = logging.getLogger(logger_name)
|
||||
# Use the logger's *effective* level at creation time
|
||||
effective_level = lg.getEffectiveLevel()
|
||||
_ensure_file_handler(
|
||||
logger=lg,
|
||||
filename=filename,
|
||||
level=effective_level,
|
||||
formatter=formatter,
|
||||
)
|
||||
|
||||
# Optionally apply the same handler setup to other loggers outside this module
|
||||
if apply_to:
|
||||
for name in apply_to:
|
||||
lg = logging.getLogger(name)
|
||||
# This removes any existing handlers so we don't duplicate output
|
||||
# and ensures these loggers share exactly the same handlers as base_logger.
|
||||
lg.handlers.clear()
|
||||
for h in base_logger.handlers:
|
||||
lg.addHandler(h)
|
||||
lg.setLevel(base_logger.level)
|
||||
# We've attached handlers directly to these loggers; if propagate
|
||||
# stayed True, records would bubble up to ancestor loggers and could be
|
||||
# emitted twice (here and on the parent/root). Setting False isolates them.
|
||||
lg.propagate = False
|
||||
|
||||
# Optionally capture warnings
|
||||
if capture_warnings:
|
||||
logging.captureWarnings(True)
|
||||
|
||||
|
||||
def setup_module(
|
||||
level: int | str = "INFO",
|
||||
*,
|
||||
name: str = "agentlightning",
|
||||
console: bool = True,
|
||||
color: bool | Dict[str, Any] = True,
|
||||
propagate: bool = False,
|
||||
disable_existing_loggers: bool = False,
|
||||
) -> logging.Logger:
|
||||
"""Initializes and returns the base logger for `agentlightning`.
|
||||
|
||||
This function constructs and applies a `dictConfig` configuration for the
|
||||
logger hierarchy rooted at `name`. It supports either rich console
|
||||
formatting (via `RichHandler`) or plain text formatting, based on the
|
||||
`color` argument.
|
||||
|
||||
Unlike [`setup_logging`][agentlightning.setup_logging], this function configures only a single logger namespace
|
||||
and does not attach extra handlers or submodule levels. It is primarily used
|
||||
internally by [`setup_logging`][agentlightning.setup_logging] but is also suitable for direct integration in
|
||||
custom logging workflows.
|
||||
"""
|
||||
root_cfg: Dict[str, Any] = {
|
||||
"version": 1,
|
||||
"disable_existing_loggers": disable_existing_loggers,
|
||||
"loggers": {
|
||||
name: {
|
||||
"handlers": [],
|
||||
"level": level,
|
||||
"propagate": propagate,
|
||||
}
|
||||
},
|
||||
"handlers": {},
|
||||
"formatters": {},
|
||||
}
|
||||
|
||||
# Choose formatter / handler definition
|
||||
if color is not False and console:
|
||||
# Console must be true to display colored outputs
|
||||
if isinstance(color, dict):
|
||||
rich_handler_config = color
|
||||
else:
|
||||
rich_handler_config: Dict[str, Any] = {
|
||||
"rich_tracebacks": False,
|
||||
"markup": False,
|
||||
"show_time": True,
|
||||
"show_path": True,
|
||||
}
|
||||
|
||||
if not _has_width():
|
||||
# e.g., in a CI environment.
|
||||
rich_handler_config["console"] = Console(width=200)
|
||||
|
||||
root_cfg["handlers"]["console"] = {
|
||||
"class": "rich.logging.RichHandler",
|
||||
"level": level,
|
||||
**rich_handler_config,
|
||||
}
|
||||
# RichHandler manages its own style; keep formatter None
|
||||
else:
|
||||
fmt_name = "plain"
|
||||
root_cfg["formatters"][fmt_name] = {
|
||||
"format": DEFAULT_FORMAT,
|
||||
"datefmt": DATE_FORMAT,
|
||||
}
|
||||
|
||||
if console:
|
||||
root_cfg["handlers"]["console"] = {
|
||||
"class": "logging.StreamHandler",
|
||||
"level": level,
|
||||
"formatter": fmt_name,
|
||||
}
|
||||
|
||||
# Attach selected handlers to agentlightning
|
||||
handler_names = list(root_cfg["handlers"].keys())
|
||||
root_cfg["loggers"][name]["handlers"] = handler_names
|
||||
|
||||
# Apply dictConfig (this resets the logger handlers)
|
||||
dictConfig(root_cfg)
|
||||
|
||||
return logging.getLogger(name)
|
||||
|
||||
|
||||
def _has_width() -> bool:
|
||||
"""Automatically determine whether the terminal has a width."""
|
||||
return sys.stdout.isatty()
|
||||
|
||||
+137
-33
@@ -11,8 +11,21 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, List, Literal, Optional, Sequence, TypeVar, cast
|
||||
from contextlib import suppress
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Sequence,
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
@@ -30,6 +43,7 @@ from agentlightning.types import (
|
||||
RolloutRawResult,
|
||||
Span,
|
||||
)
|
||||
from agentlightning.utils.system_snapshot import system_snapshot
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentlightning.execution.events import ExecutionEvent
|
||||
@@ -52,7 +66,14 @@ class LitAgentRunner(Runner[T_task]):
|
||||
worker_id: Identifier for the active worker process, if any.
|
||||
"""
|
||||
|
||||
def __init__(self, tracer: Tracer, max_rollouts: Optional[int] = None, poll_interval: float = 5.0) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
tracer: Tracer,
|
||||
max_rollouts: Optional[int] = None,
|
||||
poll_interval: float = 5.0,
|
||||
heartbeat_interval: float = 10.0,
|
||||
heartbeat_launch_mode: Literal["asyncio", "thread"] = "asyncio",
|
||||
) -> None:
|
||||
"""Initialize the agent runner.
|
||||
|
||||
Args:
|
||||
@@ -60,11 +81,16 @@ class LitAgentRunner(Runner[T_task]):
|
||||
max_rollouts: Optional cap on iterations processed by
|
||||
[`iter`][agentlightning.LitAgentRunner.iter].
|
||||
poll_interval: Seconds to wait between store polls when no work is available.
|
||||
heartbeat_interval: Seconds to wait between sending heartbeats to the store.
|
||||
heartbeat_launch_mode: Launch mode for the heartbeat loop. Can be "asyncio" or "thread".
|
||||
"asyncio" is the default and recommended mode. Use "thread" if you are experiencing blocking coroutines.
|
||||
"""
|
||||
super().__init__()
|
||||
self._tracer = tracer
|
||||
self._max_rollouts = max_rollouts
|
||||
self._poll_interval = poll_interval
|
||||
self._heartbeat_interval = heartbeat_interval
|
||||
self._heartbeat_launch_mode = heartbeat_launch_mode
|
||||
|
||||
# Set later
|
||||
self._agent: Optional[LitAgent[T_task]] = None
|
||||
@@ -253,8 +279,9 @@ class LitAgentRunner(Runner[T_task]):
|
||||
if isinstance(raw_result, float):
|
||||
# Preserve the existing spans before another span is emitted
|
||||
trace_spans = list(self._tracer.get_last_trace())
|
||||
# This will emit another span to the tracer
|
||||
reward_span = emit_reward(raw_result)
|
||||
# This will NOT emit another span to the tracer
|
||||
reward_span = emit_reward(raw_result, auto_export=False)
|
||||
# We add it to the store manually
|
||||
await store.add_otel_span(rollout.rollout_id, rollout.attempt.attempt_id, reward_span)
|
||||
trace_spans.append(reward_span)
|
||||
|
||||
@@ -304,6 +331,67 @@ class LitAgentRunner(Runner[T_task]):
|
||||
|
||||
return trace_spans
|
||||
|
||||
async def _emit_heartbeat(self, store: LightningStore) -> None:
|
||||
"""Send a heartbeat tick to the store."""
|
||||
worker_id = self.get_worker_id()
|
||||
|
||||
try:
|
||||
await store.update_worker(worker_id, system_snapshot())
|
||||
except asyncio.CancelledError:
|
||||
# bypass the exception
|
||||
raise
|
||||
except Exception:
|
||||
logger.exception("%s Unable to update worker heartbeat.", self._log_prefix())
|
||||
|
||||
def _start_heartbeat_loop(self, store: LightningStore) -> Optional[Callable[[], Awaitable[None]]]:
|
||||
"""Start a background heartbeat loop and return an async stopper."""
|
||||
|
||||
if self._heartbeat_interval <= 0:
|
||||
return None
|
||||
|
||||
if self.worker_id is None:
|
||||
logger.warning("%s Cannot start heartbeat loop without worker_id.", self._log_prefix())
|
||||
return None
|
||||
|
||||
if self._heartbeat_launch_mode == "asyncio":
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
async def heartbeat_loop() -> None:
|
||||
while not stop_event.is_set():
|
||||
await self._emit_heartbeat(store)
|
||||
with suppress(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(stop_event.wait(), timeout=self._heartbeat_interval)
|
||||
|
||||
task = asyncio.create_task(heartbeat_loop(), name=f"{self.get_worker_id()}-heartbeat")
|
||||
|
||||
async def stop() -> None:
|
||||
stop_event.set()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
return stop
|
||||
|
||||
if self._heartbeat_launch_mode == "thread":
|
||||
stop_evt = threading.Event()
|
||||
|
||||
def thread_worker() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
while not stop_evt.is_set():
|
||||
loop.run_until_complete(self._emit_heartbeat(store))
|
||||
stop_evt.wait(self._heartbeat_interval)
|
||||
|
||||
thread = threading.Thread(target=thread_worker, name=f"{self.get_worker_id()}-heartbeat", daemon=True)
|
||||
thread.start()
|
||||
|
||||
async def stop() -> None:
|
||||
stop_evt.set()
|
||||
await asyncio.to_thread(thread.join)
|
||||
|
||||
return stop
|
||||
|
||||
raise ValueError(f"Unsupported heartbeat launch mode: {self._heartbeat_launch_mode}")
|
||||
|
||||
async def _sleep_until_next_poll(self, event: Optional[ExecutionEvent] = None) -> None:
|
||||
"""Sleep until the next poll interval, with optional event-based interruption.
|
||||
|
||||
@@ -450,39 +538,49 @@ class LitAgentRunner(Runner[T_task]):
|
||||
logger.info(f"{self._log_prefix()} Started async rollouts (max: {self._max_rollouts or 'unlimited'}).")
|
||||
store = self.get_store()
|
||||
|
||||
while not (event is not None and event.is_set()) and (
|
||||
self._max_rollouts is None or num_tasks_processed < self._max_rollouts
|
||||
):
|
||||
# Retrieve the next rollout
|
||||
next_rollout: Optional[Rollout] = None
|
||||
while not (event is not None and event.is_set()):
|
||||
logger.debug(f"{self._log_prefix()} Try to poll for next rollout.")
|
||||
next_rollout = await store.dequeue_rollout()
|
||||
stop_heartbeat = self._start_heartbeat_loop(store)
|
||||
|
||||
try:
|
||||
while not (event is not None and event.is_set()) and (
|
||||
self._max_rollouts is None or num_tasks_processed < self._max_rollouts
|
||||
):
|
||||
# Retrieve the next rollout
|
||||
next_rollout: Optional[Rollout] = None
|
||||
while not (event is not None and event.is_set()):
|
||||
logger.debug(f"{self._log_prefix()} Try to poll for next rollout.")
|
||||
next_rollout = await store.dequeue_rollout(worker_id=self.get_worker_id())
|
||||
if next_rollout is None:
|
||||
logger.debug(
|
||||
f"{self._log_prefix()} No rollout to poll. Waiting for {self._poll_interval} seconds."
|
||||
)
|
||||
await self._sleep_until_next_poll(event)
|
||||
else:
|
||||
break
|
||||
|
||||
if next_rollout is None:
|
||||
logger.debug(f"{self._log_prefix()} No rollout to poll. Waiting for {self._poll_interval} seconds.")
|
||||
await self._sleep_until_next_poll(event)
|
||||
else:
|
||||
break
|
||||
return
|
||||
|
||||
if next_rollout is None:
|
||||
return
|
||||
try:
|
||||
# Claim the rollout but updating the current worker id
|
||||
await store.update_attempt(
|
||||
next_rollout.rollout_id, next_rollout.attempt.attempt_id, worker_id=self.get_worker_id()
|
||||
)
|
||||
except Exception:
|
||||
# This exception could happen if the rollout is dequeued and the other end died for some reason
|
||||
logger.exception(f"{self._log_prefix()} Exception during update_attempt, giving up the rollout.")
|
||||
continue
|
||||
|
||||
try:
|
||||
# Claim the rollout but updating the current worker id
|
||||
await store.update_attempt(
|
||||
next_rollout.rollout_id, next_rollout.attempt.attempt_id, worker_id=self.get_worker_id()
|
||||
)
|
||||
except Exception:
|
||||
# This exception could happen if the rollout is dequeued and the other end died for some reason
|
||||
logger.exception(f"{self._log_prefix()} Exception during update_attempt, giving up the rollout.")
|
||||
continue
|
||||
# Execute the step
|
||||
await self._step_impl(next_rollout)
|
||||
|
||||
# Execute the step
|
||||
await self._step_impl(next_rollout)
|
||||
|
||||
num_tasks_processed += 1
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(f"{self._log_prefix()} Progress: {num_tasks_processed}/{self._max_rollouts or 'unlimited'}")
|
||||
num_tasks_processed += 1
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(
|
||||
f"{self._log_prefix()} Progress: {num_tasks_processed}/{self._max_rollouts or 'unlimited'}"
|
||||
)
|
||||
finally:
|
||||
if stop_heartbeat is not None:
|
||||
await stop_heartbeat()
|
||||
|
||||
logger.info(f"{self._log_prefix()} Finished async rollouts. Processed {num_tasks_processed} tasks.")
|
||||
|
||||
@@ -526,6 +624,12 @@ class LitAgentRunner(Runner[T_task]):
|
||||
resources_id = None
|
||||
|
||||
attempted_rollout = await self.get_store().start_rollout(input=input, mode=mode, resources_id=resources_id)
|
||||
# Register the attempt as running by the current worker
|
||||
await self.get_store().update_attempt(
|
||||
attempted_rollout.rollout_id,
|
||||
attempted_rollout.attempt.attempt_id,
|
||||
worker_id=self.get_worker_id(),
|
||||
)
|
||||
rollout_id = await self._step_impl(attempted_rollout, raise_on_exception=True)
|
||||
|
||||
completed_rollout = await store.get_rollout_by_id(rollout_id)
|
||||
|
||||
@@ -142,7 +142,13 @@ class ServerDataStore:
|
||||
async with self._resources_lock:
|
||||
resources = self._resource_versions.get(resources_id)
|
||||
if resources:
|
||||
return ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
return ResourcesUpdate(
|
||||
resources_id=resources_id,
|
||||
resources=resources,
|
||||
create_time=time.time(),
|
||||
update_time=time.time(),
|
||||
version=1,
|
||||
)
|
||||
return None
|
||||
|
||||
async def get_latest_resources(self) -> Optional[ResourcesUpdate]:
|
||||
@@ -357,7 +363,9 @@ class AgentLightningServer:
|
||||
if not self._store:
|
||||
raise RuntimeError("Store not initialized. The server may not be running.")
|
||||
resources_id = f"res-{uuid.uuid4()}"
|
||||
update = ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
update = ResourcesUpdate(
|
||||
resources_id=resources_id, resources=resources, create_time=time.time(), update_time=time.time(), version=1
|
||||
)
|
||||
await self._store.update_resources(update)
|
||||
return resources_id
|
||||
|
||||
|
||||
@@ -1,14 +1,17 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from .base import LightningStore
|
||||
from .base import LightningStore, LightningStoreCapabilities
|
||||
from .client_server import LightningStoreClient, LightningStoreServer
|
||||
from .collection_based import CollectionBasedLightningStore
|
||||
from .memory import InMemoryLightningStore
|
||||
from .threading import LightningStoreThreaded
|
||||
|
||||
__all__ = [
|
||||
"LightningStore",
|
||||
"LightningStoreCapabilities",
|
||||
"LightningStoreClient",
|
||||
"LightningStoreServer",
|
||||
"InMemoryLightningStore",
|
||||
"CollectionBasedLightningStore",
|
||||
"LightningStoreThreaded",
|
||||
]
|
||||
|
||||
+249
-12
@@ -2,7 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Literal, Optional, Sequence
|
||||
from typing import Any, Dict, List, Literal, Optional, Sequence, TypedDict
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
@@ -17,6 +17,8 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
WorkerStatus,
|
||||
)
|
||||
|
||||
|
||||
@@ -52,6 +54,22 @@ UNSET = _UnsetType()
|
||||
Unset = _UnsetType # Alias for convenience
|
||||
|
||||
|
||||
class LightningStoreCapabilities(TypedDict, total=False):
|
||||
"""Capability of a LightningStore implementation.
|
||||
|
||||
All keys are optional and false by default.
|
||||
"""
|
||||
|
||||
thread_safe: bool
|
||||
"""Whether the store is thread-safe."""
|
||||
async_safe: bool
|
||||
"""Whether the store is async-safe."""
|
||||
zero_copy: bool
|
||||
"""Whether the store has only one copy across all threads/processes."""
|
||||
otlp_traces: bool
|
||||
"""Whether the store supports OTLP/HTTP traces."""
|
||||
|
||||
|
||||
class LightningStore:
|
||||
"""Contract for the persistent control-plane that coordinates training rollouts.
|
||||
|
||||
@@ -74,6 +92,32 @@ class LightningStore:
|
||||
Unless stated otherwise, missing identifiers should result in a `ValueError`.
|
||||
"""
|
||||
|
||||
@property
|
||||
def capabilities(self) -> LightningStoreCapabilities:
|
||||
"""Return the capabilities of the store."""
|
||||
return LightningStoreCapabilities(
|
||||
thread_safe=False,
|
||||
async_safe=False,
|
||||
zero_copy=False,
|
||||
otlp_traces=False,
|
||||
)
|
||||
|
||||
def otlp_traces_endpoint(self) -> str:
|
||||
"""Return the OTLP/HTTP traces endpoint of the store.
|
||||
|
||||
The traces can have rollout ID and attempt ID (and optionally sequence ID)
|
||||
saved in the "resource" of the spans.
|
||||
The store, if it supports OTLP, should be able to receive the traces and save them
|
||||
via [`add_span`][agentlightning.LightningStore.add_span] or
|
||||
[`add_otel_span`][agentlightning.LightningStore.add_otel_span].
|
||||
|
||||
The endpoint should be compatible with [OTLP HTTP protocol](https://opentelemetry.io/docs/specs/otlp/).
|
||||
It's not necessarily compatible with OTLP gRPC protocol.
|
||||
|
||||
The returned endpoint will usually ends with `/v1/traces`.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def start_rollout(
|
||||
self,
|
||||
input: TaskInput,
|
||||
@@ -148,7 +192,7 @@ class LightningStore:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
"""Claim the oldest queued rollout and transition it to `preparing`.
|
||||
|
||||
This function do not block.
|
||||
@@ -161,6 +205,8 @@ class LightningStore:
|
||||
the number of attempts already registered for the rollout plus one.
|
||||
* Return an [`AttemptedRollout`][agentlightning.AttemptedRollout] snapshot so the
|
||||
runner knows both rollout metadata and the attempt identifier.
|
||||
* Optionally refresh the caller's [`Worker`][agentlightning.Worker] telemetry
|
||||
(e.g., `last_dequeue_time`) when `worker_id` is provided.
|
||||
|
||||
Returns:
|
||||
The next attempt to execute, or `None` when no eligible rollouts are queued.
|
||||
@@ -247,30 +293,77 @@ class LightningStore:
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_rollouts(
|
||||
self, *, status: Optional[Sequence[RolloutStatus]] = None, rollout_ids: Optional[Sequence[str]] = None
|
||||
) -> List[Rollout]:
|
||||
self,
|
||||
*,
|
||||
status_in: Optional[Sequence[RolloutStatus]] = None,
|
||||
rollout_id_in: Optional[Sequence[str]] = None,
|
||||
rollout_id_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
# Deprecated fields
|
||||
status: Optional[Sequence[RolloutStatus]] = None,
|
||||
rollout_ids: Optional[Sequence[str]] = None,
|
||||
) -> Sequence[Rollout]:
|
||||
"""Retrieve rollouts filtered by status and/or explicit identifiers.
|
||||
|
||||
This interface supports structured filtering, sorting, and pagination so
|
||||
callers can build simple dashboards without copying data out of the
|
||||
store. The legacy parameters `status` and `rollout_ids` remain valid and
|
||||
are treated as aliases for `status_in` and `rollout_id_in`
|
||||
respectively—when both the new and deprecated parameters are supplied
|
||||
the new parameters take precedence.
|
||||
|
||||
Args:
|
||||
status: Optional whitelist of [`RolloutStatus`][agentlightning.RolloutStatus] values.
|
||||
rollout_ids: Optional whitelist of rollout identifiers to include.
|
||||
status_in: Optional whitelist of [`RolloutStatus`][agentlightning.RolloutStatus] values.
|
||||
rollout_id_in: Optional whitelist of rollout identifiers to include.
|
||||
rollout_id_contains: Optional substring match for rollout identifiers.
|
||||
filter_logic: Logical operator to combine filters.
|
||||
sort_by: Optional field to sort by. Must reference a numeric or string
|
||||
field on [`Rollout`][agentlightning.Rollout].
|
||||
sort_order: Direction to sort when `sort_by` is provided.
|
||||
limit: Maximum number of rows to return. Use `-1` for "no limit".
|
||||
offset: Number of rows to skip before returning results.
|
||||
status: Deprecated field. Use `status_in` instead.
|
||||
rollout_ids: Deprecated field. Use `rollout_id_in` instead.
|
||||
|
||||
Returns:
|
||||
A list of matching rollouts. Ordering is backend-defined but must be deterministic.
|
||||
A sequence of matching rollouts (or [`AttemptedRollout`][agentlightning.AttemptedRollout]
|
||||
when attempts exist). Ordering is deterministic when `sort_by` is set.
|
||||
The return value is not guaranteed to be a list.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement the query.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_attempts(self, rollout_id: str) -> List[Attempt]:
|
||||
async def query_attempts(
|
||||
self,
|
||||
rollout_id: str,
|
||||
*,
|
||||
sort_by: Optional[str] = "sequence_id",
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[Attempt]:
|
||||
"""Return every attempt ever created for `rollout_id` in ascending sequence order.
|
||||
|
||||
The parameters allow callers to re-order or paginate the attempts so that
|
||||
large retry histories can be streamed lazily.
|
||||
|
||||
Args:
|
||||
rollout_id: Identifier of the rollout being inspected.
|
||||
sort_by: Field to sort by. Must be a numeric or string field of
|
||||
[`Attempt`][agentlightning.Attempt]. Defaults to `sequence_id` (oldest first).
|
||||
sort_order: Order to sort by.
|
||||
limit: Limit on the number of results. `-1` for unlimited.
|
||||
offset: Offset into the results.
|
||||
|
||||
Returns:
|
||||
Attempts sorted by `sequence_id` (oldest first). Returns an empty list when none exist.
|
||||
Sequence of Attempts. Returns an empty sequence when none exist.
|
||||
The return value is not guaranteed to be a list.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement the query.
|
||||
@@ -307,6 +400,41 @@ class LightningStore:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_resources(
|
||||
self,
|
||||
*,
|
||||
resources_id: Optional[str] = None,
|
||||
resources_id_contains: Optional[str] = None,
|
||||
# Filter logic is not supported here because I can't see why it's needed.
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[ResourcesUpdate]:
|
||||
"""List every stored resource snapshot in insertion order.
|
||||
|
||||
Supports lightweight filtering, sorting, and pagination for embedding in
|
||||
dashboards.
|
||||
|
||||
Args:
|
||||
resources_id: Optional identifier of the resources to include.
|
||||
resources_id_contains: Optional substring match for resources identifiers.
|
||||
sort_by: Optional field to sort by (must be numeric or string on
|
||||
[`ResourcesUpdate`][agentlightning.ResourcesUpdate]).
|
||||
sort_order: Order to sort by.
|
||||
limit: Limit on the number of results. `-1` for unlimited.
|
||||
offset: Offset into the results.
|
||||
|
||||
Returns:
|
||||
[`ResourcesUpdate`][agentlightning.ResourcesUpdate] objects.
|
||||
By default, resources are sorted in a deterministic but undefined order.
|
||||
The return value is not guaranteed to be a list.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement retrieval.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def get_resources_by_id(self, resources_id: str) -> Optional[ResourcesUpdate]:
|
||||
"""Return a specific named resource snapshot by identifier.
|
||||
|
||||
@@ -383,19 +511,61 @@ class LightningStore:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_spans(self, rollout_id: str, attempt_id: str | Literal["latest"] | None = None) -> List[Span]:
|
||||
async def query_spans(
|
||||
self,
|
||||
rollout_id: str,
|
||||
attempt_id: str | Literal["latest"] | None = None,
|
||||
*,
|
||||
# Filtering
|
||||
trace_id: Optional[str] = None,
|
||||
trace_id_contains: Optional[str] = None,
|
||||
span_id: Optional[str] = None,
|
||||
span_id_contains: Optional[str] = None,
|
||||
parent_id: Optional[str] = None,
|
||||
parent_id_contains: Optional[str] = None,
|
||||
name: Optional[str] = None,
|
||||
name_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
# Pagination
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
# Sorting
|
||||
sort_by: Optional[str] = "sequence_id",
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
) -> Sequence[Span]:
|
||||
"""Return the stored spans for a rollout, optionally scoped to one attempt.
|
||||
|
||||
Spans must be returned in ascending `sequence_id` order. Implementations may raise
|
||||
a `RuntimeError` when spans were evicted or expired.
|
||||
Supports a handful of filters that cover the most common debugging
|
||||
scenarios (matching `trace_id`/`span_id`/`parent_id` or substring
|
||||
matches on the span name). `attempt_id="latest"` acts as a convenience
|
||||
that resolves the most recent attempt before evaluating filters. When
|
||||
`attempt_id=None`, spans across every attempt are eligible. By default
|
||||
results are sorted by `sequence_id` (oldest first). Implementations may
|
||||
raise a `RuntimeError` when spans were evicted or expired.
|
||||
|
||||
Args:
|
||||
rollout_id: Identifier of the rollout being inspected.
|
||||
attempt_id: Attempt identifier to filter by. Pass `"latest"` to retrieve only the
|
||||
most recent attempt, or `None` to return all spans across attempts.
|
||||
trace_id: Optional trace ID to filter by.
|
||||
trace_id_contains: Optional substring match for trace IDs.
|
||||
span_id: Optional span ID to filter by.
|
||||
span_id_contains: Optional substring match for span IDs.
|
||||
parent_id: Optional parent span ID to filter by.
|
||||
parent_id_contains: Optional substring match for parent span IDs.
|
||||
name: Optional span name to filter by.
|
||||
name_contains: Optional substring match for span names.
|
||||
filter_logic: Logical operator to combine the optional filters above.
|
||||
The `rollout_id` argument is always applied with AND semantics.
|
||||
limit: Limit on the number of results. `-1` for unlimited.
|
||||
offset: Offset into the results.
|
||||
sort_by: Field to sort by. Must be a numeric or string field of
|
||||
[`Span`][agentlightning.Span].
|
||||
sort_order: Order to sort by.
|
||||
|
||||
Returns:
|
||||
An ordered list of spans (possibly empty).
|
||||
The return value is not guaranteed to be a list.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement the query.
|
||||
@@ -497,6 +667,12 @@ class LightningStore:
|
||||
Similar to [`update_rollout()`][agentlightning.LightningStore.update_rollout],
|
||||
parameters also default to the sentinel [`UNSET`][agentlightning.store.base.UNSET].
|
||||
|
||||
If `worker_id` is present, the worker status will be updated following the rules:
|
||||
|
||||
1. If attempt status is "succeeded" or "failed", the corresponding worker status will be set to "idle".
|
||||
2. If attempt status is "unresponsive" or "timeout", the corresponding worker status will be set to "unknown".
|
||||
3. Otherwise, the worker status will be set to "busy".
|
||||
|
||||
Args:
|
||||
rollout_id: Identifier of the rollout whose attempt will be updated.
|
||||
attempt_id: Attempt identifier or `"latest"` as a convenience.
|
||||
@@ -513,3 +689,64 @@ class LightningStore:
|
||||
ValueError: Implementations must raise when the rollout or attempt is unknown.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query_workers(
|
||||
self,
|
||||
*,
|
||||
status_in: Optional[Sequence[WorkerStatus]] = None,
|
||||
worker_id_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[Worker]:
|
||||
"""Query all workers in the system.
|
||||
|
||||
Args:
|
||||
status_in: Optional whitelist of [`WorkerStatus`][agentlightning.WorkerStatus] values.
|
||||
worker_id_contains: Optional substring match for worker identifiers.
|
||||
filter_logic: Logical operator to combine the optional filters above.
|
||||
sort_by: Field to sort by. Must be a numeric or string field of [`Worker`][agentlightning.Worker].
|
||||
sort_order: Order to sort by.
|
||||
limit: Limit on the number of results. `-1` for unlimited.
|
||||
offset: Offset into the results.
|
||||
|
||||
Returns:
|
||||
Sequence of Workers. Returns an empty sequence when none exist.
|
||||
The return value is not guaranteed to be a list.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
"""Retrieve a single worker by identifier.
|
||||
|
||||
Args:
|
||||
worker_id: Identifier of the worker.
|
||||
|
||||
Returns:
|
||||
The worker record if it exists, otherwise `None`.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Subclasses must implement lookup semantics.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
"""Record a heartbeat for `worker_id` and refresh telemetry.
|
||||
|
||||
Implementations must treat this API as heartbeat-only: it should snapshot
|
||||
the latest stats when provided, stamp `last_heartbeat_time` with the
|
||||
current wall clock, and rely on other store mutations (`dequeue_rollout`,
|
||||
`update_attempt`, etc.) to drive the worker's busy/idle status,
|
||||
assignment, and activity timestamps.
|
||||
|
||||
Args:
|
||||
worker_id: Identifier of the worker to update.
|
||||
heartbeat_stats: Replacement worker heartbeat statistics (non-null when provided).
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,18 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from .base import Collection, FilterOptions, KeyValue, LightningCollections, PaginatedResult, Queue, SortOptions
|
||||
from .memory import DequeBasedQueue, DictBasedKeyValue, InMemoryLightningCollections, ListBasedCollection
|
||||
|
||||
__all__ = [
|
||||
"Collection",
|
||||
"Queue",
|
||||
"KeyValue",
|
||||
"FilterOptions",
|
||||
"SortOptions",
|
||||
"PaginatedResult",
|
||||
"LightningCollections",
|
||||
"ListBasedCollection",
|
||||
"DequeBasedQueue",
|
||||
"DictBasedKeyValue",
|
||||
"InMemoryLightningCollections",
|
||||
]
|
||||
@@ -0,0 +1,263 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncContextManager,
|
||||
Generic,
|
||||
Optional,
|
||||
Sequence,
|
||||
Type,
|
||||
TypeVar,
|
||||
)
|
||||
|
||||
from agentlightning.types import (
|
||||
Attempt,
|
||||
FilterOptions,
|
||||
PaginatedResult,
|
||||
ResourcesUpdate,
|
||||
Rollout,
|
||||
SortOptions,
|
||||
Span,
|
||||
Worker,
|
||||
)
|
||||
|
||||
T = TypeVar("T") # Recommended to be a BaseModel
|
||||
K = TypeVar("K")
|
||||
V = TypeVar("V")
|
||||
|
||||
|
||||
class Collection(Generic[T]):
|
||||
"""Behaves like a list of items. Supporting addition, updating, and deletion of items."""
|
||||
|
||||
def primary_keys(self) -> Sequence[str]:
|
||||
"""Get the primary keys of the collection."""
|
||||
raise NotImplementedError()
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<{self.__class__.__name__}[{self.item_type().__name__}] ({self.size()})>"
|
||||
|
||||
def item_type(self) -> Type[T]:
|
||||
"""Get the type of the items in the collection."""
|
||||
raise NotImplementedError()
|
||||
|
||||
def size(self) -> int:
|
||||
"""Get the number of items in the collection."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def query(
|
||||
self,
|
||||
filter: Optional[FilterOptions] = None,
|
||||
sort: Optional[SortOptions] = None,
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> PaginatedResult[T]:
|
||||
"""Query the collection with the given filters, sort order, and pagination.
|
||||
|
||||
Args:
|
||||
filter:
|
||||
The filters to apply to the collection. See [`FilterOptions`][agentlightning.FilterOptions].
|
||||
|
||||
sort:
|
||||
The options for sorting the collection. See [`SortOptions`][agentlightning.SortOptions].
|
||||
The field must exist in the model. If field might contain null values, in which case the behavior is undefined
|
||||
(i.e., depending on the implementation).
|
||||
|
||||
limit:
|
||||
Max number of items to return. Use -1 for "no limit".
|
||||
|
||||
offset:
|
||||
Number of items to skip from the start of the *matching* items.
|
||||
|
||||
Returns:
|
||||
PaginatedResult with items, limit, offset, and total matched items.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def get(
|
||||
self,
|
||||
filter: Optional[FilterOptions] = None,
|
||||
sort: Optional[SortOptions] = None,
|
||||
) -> Optional[T]:
|
||||
"""Get the first item that matches the given filters.
|
||||
|
||||
Args:
|
||||
filter: The filters to apply to the collection.
|
||||
See [`FilterOptions`][agentlightning.store.collection.FilterOptions].
|
||||
sort: Sort options. See [`SortOptions`][agentlightning.store.collection.SortOptions].
|
||||
|
||||
Returns:
|
||||
The first item that matches the given filters, or None if no item matches.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def insert(self, items: Sequence[T]) -> None:
|
||||
"""Add the given items to the collection.
|
||||
|
||||
Raises:
|
||||
ValueError: If an item with the same primary key already exists.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def update(self, items: Sequence[T]) -> None:
|
||||
"""Update the given items in the collection.
|
||||
|
||||
Raises:
|
||||
ValueError: If an item with the primary keys does not exist.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def upsert(self, items: Sequence[T]) -> None:
|
||||
"""Upsert the given items into the collection.
|
||||
|
||||
If the items with the same primary keys already exist, they will be updated.
|
||||
Otherwise, they will be inserted.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def delete(self, items: Sequence[T]) -> None:
|
||||
"""Delete the given items from the collection.
|
||||
|
||||
Args:
|
||||
items: The items to delete from the collection.
|
||||
|
||||
Raises:
|
||||
ValueError: If the items with the primary keys to be deleted do not exist.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class Queue(Generic[T]):
|
||||
"""Behaves like a deque. Supporting appending items to the end and popping items from the front."""
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<{self.__class__.__name__}[{self.item_type().__name__}] ({self.size()})>"
|
||||
|
||||
def item_type(self) -> Type[T]:
|
||||
"""Get the type of the items in the queue."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def has(self, item: T) -> bool:
|
||||
"""Check if the given item is in the queue."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def enqueue(self, items: Sequence[T]) -> Sequence[T]:
|
||||
"""Append the given items to the end of the queue.
|
||||
|
||||
Args:
|
||||
items: The items to append to the end of the queue.
|
||||
|
||||
Returns:
|
||||
The items that were appended to the end of the queue.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def dequeue(self, limit: int = 1) -> Sequence[T]:
|
||||
"""Pop the given number of items from the front of the queue.
|
||||
|
||||
Args:
|
||||
limit: The number of items to pop from the front of the queue.
|
||||
|
||||
Returns:
|
||||
The items that were popped from the front of the queue.
|
||||
If there are less than `limit` items in the queue, the remaining items will be returned.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def peek(self, limit: int = 1) -> Sequence[T]:
|
||||
"""Peek the given number of items from the front of the queue.
|
||||
|
||||
Args:
|
||||
limit: The number of items to peek from the front of the queue.
|
||||
|
||||
Returns:
|
||||
The items that were peeked from the front of the queue.
|
||||
If there are less than `limit` items in the queue, the remaining items will be returned.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def size(self) -> int:
|
||||
"""Get the number of items in the queue."""
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class KeyValue(Generic[K, V]):
|
||||
"""Behaves like a dictionary. Supporting addition, updating, and deletion of items."""
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<{self.__class__.__name__} ({self.size()})>"
|
||||
|
||||
async def has(self, key: K) -> bool:
|
||||
"""Check if the given key is in the dictionary."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def get(self, key: K, default: V | None = None) -> V | None:
|
||||
"""Get the value for the given key, or the default value if the key is not found."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def set(self, key: K, value: V) -> None:
|
||||
"""Set the value for the given key."""
|
||||
raise NotImplementedError()
|
||||
|
||||
async def pop(self, key: K, default: V | None = None) -> V | None:
|
||||
"""Pop the value for the given key, or the default value if the key is not found."""
|
||||
raise NotImplementedError()
|
||||
|
||||
def size(self) -> int:
|
||||
"""Get the number of items in the dictionary."""
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class LightningCollections:
|
||||
"""Collections of rollouts, attempts, spans, resources, and workers.
|
||||
|
||||
[LightningStore][agentlightning.LightningStore] implementations can use this as a storage base
|
||||
to implement the store API.
|
||||
"""
|
||||
|
||||
@property
|
||||
def rollouts(self) -> Collection[Rollout]:
|
||||
"""Collections of rollouts."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def attempts(self) -> Collection[Attempt]:
|
||||
"""Collections of attempts."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def spans(self) -> Collection[Span]:
|
||||
"""Collections of spans."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def resources(self) -> Collection[ResourcesUpdate]:
|
||||
"""Collections of resources."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def workers(self) -> Collection[Worker]:
|
||||
"""Collections of workers."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def rollout_queue(self) -> Queue[str]:
|
||||
"""Queue of rollouts (tasks)."""
|
||||
raise NotImplementedError()
|
||||
|
||||
@property
|
||||
def span_sequence_ids(self) -> KeyValue[str, int]:
|
||||
"""Dictionary (counter) of span sequence IDs."""
|
||||
raise NotImplementedError()
|
||||
|
||||
def atomic(self, *args: Any, **kwargs: Any) -> AsyncContextManager[None]:
|
||||
"""Perform a atomic operation on the collections.
|
||||
|
||||
Subclass may use args and kwargs to support multiple levels of atomicity.
|
||||
|
||||
Args:
|
||||
*args: Arguments to pass to the operation.
|
||||
**kwargs: Keyword arguments to pass to the operation.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,811 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import weakref
|
||||
from collections import deque
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Deque,
|
||||
Dict,
|
||||
Iterable,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
MutableMapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from agentlightning.types import (
|
||||
Attempt,
|
||||
FilterField,
|
||||
FilterOptions,
|
||||
PaginatedResult,
|
||||
ResourcesUpdate,
|
||||
Rollout,
|
||||
SortOptions,
|
||||
Span,
|
||||
Worker,
|
||||
)
|
||||
|
||||
from .base import (
|
||||
Collection,
|
||||
KeyValue,
|
||||
LightningCollections,
|
||||
Queue,
|
||||
)
|
||||
|
||||
T = TypeVar("T") # Recommended to be a BaseModel, not a dict
|
||||
K = TypeVar("K")
|
||||
V = TypeVar("V")
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Nested structure type:
|
||||
# dict[pk1] -> dict[pk2] -> ... -> item
|
||||
ListBasedCollectionItemType = Union[
|
||||
Dict[Any, "ListBasedCollectionItemType[T]"], # intermediate node
|
||||
Dict[Any, T], # leaf node dictionary
|
||||
]
|
||||
|
||||
FilterMap = Mapping[str, FilterField]
|
||||
MutationMode = Literal["insert", "update", "upsert", "delete"]
|
||||
|
||||
|
||||
def _merge_must_filters(target: Dict[str, FilterField], definition: Any) -> None:
|
||||
"""Normalize a `_must` filter group into the provided mapping.
|
||||
|
||||
Mainly for validation purposes.
|
||||
"""
|
||||
if definition is None:
|
||||
return
|
||||
|
||||
entries: List[Mapping[str, FilterField]] = []
|
||||
if isinstance(definition, Mapping):
|
||||
entries.append(cast(Mapping[str, FilterField], definition))
|
||||
elif isinstance(definition, Sequence) and not isinstance(definition, (str, bytes)):
|
||||
for entry in definition: # type: ignore
|
||||
if not isinstance(entry, Mapping):
|
||||
raise TypeError("Each `_must` entry must be a mapping of field names to operators")
|
||||
entries.append(cast(Mapping[str, FilterField], entry))
|
||||
else:
|
||||
raise TypeError("`_must` filters must be provided as a mapping or sequence of mappings")
|
||||
|
||||
for entry in entries:
|
||||
for field_name, ops in entry.items():
|
||||
existing = target.get(field_name, {})
|
||||
merged_ops: Dict[str, Any] = dict(existing)
|
||||
for op_name, expected in ops.items():
|
||||
if op_name in merged_ops:
|
||||
raise ValueError(f"Duplicate operator '{op_name}' for field '{field_name}' in must filters")
|
||||
merged_ops[op_name] = expected
|
||||
target[field_name] = cast(FilterField, merged_ops)
|
||||
|
||||
|
||||
def _normalize_filter_options(
|
||||
filter_options: Optional[FilterOptions],
|
||||
) -> Tuple[Optional[FilterMap], Optional[FilterMap], Literal["and", "or"]]:
|
||||
"""Convert FilterOptions to the internal structure and resolve aggregate logic."""
|
||||
if not filter_options:
|
||||
return None, None, "and"
|
||||
|
||||
aggregate = cast(Literal["and", "or"], filter_options.get("_aggregate", "and"))
|
||||
if aggregate not in ("and", "or"):
|
||||
raise ValueError(f"Unsupported filter aggregate '{aggregate}'")
|
||||
|
||||
# Extract normalized filters and must filters from the filter options.
|
||||
normalized: Dict[str, FilterField] = {}
|
||||
must_filters: Dict[str, FilterField] = {}
|
||||
for field_name, ops in filter_options.items():
|
||||
if field_name == "_aggregate":
|
||||
continue
|
||||
if field_name == "_must":
|
||||
_merge_must_filters(must_filters, ops)
|
||||
continue
|
||||
normalized[field_name] = cast(FilterField, dict(ops)) # type: ignore
|
||||
|
||||
return (normalized or None, must_filters or None, aggregate)
|
||||
|
||||
|
||||
def _resolve_sort_options(sort: Optional[SortOptions]) -> Tuple[Optional[str], Literal["asc", "desc"]]:
|
||||
"""Extract sort field/order from the caller-provided SortOptions."""
|
||||
if not sort:
|
||||
return None, "asc"
|
||||
|
||||
sort_name = sort.get("name")
|
||||
if not sort_name:
|
||||
raise ValueError("Sort options must include a 'name' field")
|
||||
|
||||
sort_order = sort.get("order", "asc")
|
||||
if sort_order not in ("asc", "desc"):
|
||||
raise ValueError(f"Unsupported sort order '{sort_order}'")
|
||||
|
||||
return sort_name, sort_order
|
||||
|
||||
|
||||
def _item_matches_filters(
|
||||
item: object,
|
||||
filters: Optional[FilterMap],
|
||||
filter_logic: Literal["and", "or"],
|
||||
must_filters: Optional[FilterMap] = None,
|
||||
) -> bool:
|
||||
"""Check whether an item matches the provided filter definition.
|
||||
|
||||
Filter format:
|
||||
|
||||
```json
|
||||
{
|
||||
"_aggregate": "or",
|
||||
"field_name": {
|
||||
"exact": <value>,
|
||||
"within": <iterable_of_allowed_values>,
|
||||
"contains": <substring_or_element>,
|
||||
},
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
Operators within the same field are stored in a unified pool and combined using
|
||||
a universal logical operator.
|
||||
"""
|
||||
if must_filters and not _item_matches_filters(item, must_filters, "and"):
|
||||
return False
|
||||
|
||||
if not filters:
|
||||
return True
|
||||
|
||||
all_conditions_match: List[bool] = []
|
||||
|
||||
for field_name, ops in filters.items():
|
||||
item_value = getattr(item, field_name, None)
|
||||
|
||||
for op_name, expected in ops.items():
|
||||
# Ignore no-op filters
|
||||
if expected is None:
|
||||
continue
|
||||
|
||||
if op_name == "exact":
|
||||
all_conditions_match.append(item_value == expected)
|
||||
|
||||
elif op_name == "within":
|
||||
try:
|
||||
all_conditions_match.append(item_value in expected) # type: ignore[arg-type]
|
||||
except TypeError:
|
||||
all_conditions_match.append(False)
|
||||
|
||||
elif op_name == "contains":
|
||||
if item_value is None:
|
||||
all_conditions_match.append(False)
|
||||
elif isinstance(item_value, str) and isinstance(expected, str):
|
||||
all_conditions_match.append(expected in item_value)
|
||||
else:
|
||||
# Fallback: treat as generic iterable containment.
|
||||
try:
|
||||
all_conditions_match.append(expected in item_value) # type: ignore[arg-type]
|
||||
except TypeError:
|
||||
all_conditions_match.append(False)
|
||||
else:
|
||||
raise ValueError(f"Unsupported filter operator '{op_name}' for field '{field_name}'")
|
||||
|
||||
return all(all_conditions_match) if filter_logic == "and" else any(all_conditions_match)
|
||||
|
||||
|
||||
def _get_sort_value(item: object, sort_by: str) -> Any:
|
||||
"""Get a sort key for the given item/field.
|
||||
|
||||
- If the field name ends with '_time', values are treated as comparable timestamps.
|
||||
- For other fields we try to infer a safe default from the Pydantic model annotation.
|
||||
"""
|
||||
value = getattr(item, sort_by, None)
|
||||
|
||||
if sort_by.endswith("_time"):
|
||||
# For *_time fields, push missing values to the end.
|
||||
return float("inf") if value is None else value
|
||||
|
||||
if value is None:
|
||||
# Introspect model field type to choose a reasonable default for None.
|
||||
model_fields = getattr(item.__class__, "model_fields", {})
|
||||
if sort_by not in model_fields:
|
||||
raise ValueError(
|
||||
f"Failed to sort items by '{sort_by}': field does not exist " f"on {item.__class__.__name__}"
|
||||
)
|
||||
|
||||
field_type_str = str(model_fields[sort_by].annotation)
|
||||
if "str" in field_type_str or "Literal" in field_type_str:
|
||||
return ""
|
||||
if "int" in field_type_str:
|
||||
return 0
|
||||
if "float" in field_type_str:
|
||||
return 0.0
|
||||
raise ValueError(f"Failed to sort items by '{sort_by}': unsupported field type {field_type_str!r}")
|
||||
|
||||
return value
|
||||
|
||||
|
||||
class ListBasedCollection(Collection[T]):
|
||||
"""In-memory implementation of Collection using a nested dict for O(1) primary-key lookup.
|
||||
|
||||
The internal structure is:
|
||||
|
||||
{
|
||||
pk1_value: {
|
||||
pk2_value: {
|
||||
...
|
||||
pkN_value: item
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
where the nesting depth equals the number of primary keys.
|
||||
|
||||
Sorting behavior:
|
||||
|
||||
1. If no sort_by is provided, the items are returned in the order of insertion.
|
||||
2. If sort_by is provided, the items are sorted by the value of the sort_by field.
|
||||
3. If the sort_by field is a timestamp, the null values are treated as infinity.
|
||||
4. If the sort_by field is not a timestamp, the null values are treated as empty string
|
||||
if the field is str-like, 0 if the field is int-like, 0.0 if the field is float-like.
|
||||
"""
|
||||
|
||||
def __init__(self, items: List[T], item_type: Type[T], primary_keys: Sequence[str]):
|
||||
if not primary_keys:
|
||||
raise ValueError("primary_keys must be non-empty")
|
||||
|
||||
self._items: Dict[Any, Any] = {}
|
||||
self._size: int = 0
|
||||
if issubclass(item_type, dict):
|
||||
raise TypeError(f"Expect item to be not a dict, got {item_type.__name__}")
|
||||
self._item_type: Type[T] = item_type
|
||||
self._primary_keys: Tuple[str, ...] = tuple(primary_keys)
|
||||
|
||||
# Pre-populate the collection with the given items.
|
||||
for item in items or []:
|
||||
self._mutate_single(item, mode="insert")
|
||||
|
||||
def primary_keys(self) -> Sequence[str]:
|
||||
"""Return the primary key field names for this collection."""
|
||||
return self._primary_keys
|
||||
|
||||
def item_type(self) -> Type[T]:
|
||||
"""Return the Pydantic model type of items stored in this collection."""
|
||||
return self._item_type
|
||||
|
||||
def size(self) -> int:
|
||||
"""Return the number of items stored in the collection."""
|
||||
return self._size
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<{self.__class__.__name__}[{self.item_type().__name__}] ({self.size()})>"
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def _ensure_item_type(self, item: T) -> None:
|
||||
"""Validate that the item matches the declared item_type."""
|
||||
if not isinstance(item, self._item_type):
|
||||
raise TypeError(f"Expected item of type {self._item_type.__name__}, " f"got {type(item).__name__}")
|
||||
|
||||
def _extract_primary_key_values(self, item: T) -> Tuple[Any, ...]:
|
||||
"""Extract the primary key values from an item.
|
||||
|
||||
Raises:
|
||||
ValueError: If any primary key is missing on the item.
|
||||
"""
|
||||
values: List[Any] = []
|
||||
for key in self._primary_keys:
|
||||
if not hasattr(item, key):
|
||||
raise ValueError(f"Item {item} does not have primary key field '{key}'")
|
||||
values.append(getattr(item, key))
|
||||
return tuple(values)
|
||||
|
||||
def _render_key_values(self, key_values: Sequence[Any]) -> str:
|
||||
return ", ".join(f"{name}={value!r}" for name, value in zip(self._primary_keys, key_values))
|
||||
|
||||
def _locate_node(
|
||||
self,
|
||||
key_values: Sequence[Any],
|
||||
create_missing: bool,
|
||||
) -> Tuple[MutableMapping[Any, Any], Any]:
|
||||
"""Locate the parent mapping and final key for an item path.
|
||||
|
||||
Args:
|
||||
key_values: The sequence of primary key values.
|
||||
create_missing: Whether to create intermediate dictionaries as needed.
|
||||
|
||||
Returns:
|
||||
(parent_mapping, final_key)
|
||||
|
||||
Raises:
|
||||
KeyError: If the path does not exist and create_missing is False.
|
||||
ValueError: If the internal structure is corrupted (non-dict where dict is expected).
|
||||
"""
|
||||
if not key_values:
|
||||
raise ValueError("key_values must be non-empty")
|
||||
|
||||
current: MutableMapping[Any, Any] = self._items
|
||||
for idx, value in enumerate(key_values):
|
||||
is_last = idx == len(key_values) - 1
|
||||
if is_last:
|
||||
# At the final level, current[value] is the item (or will be).
|
||||
return current, value # type: ignore
|
||||
|
||||
# Intermediate level: current[value] must be a dict.
|
||||
if value not in current:
|
||||
if not create_missing:
|
||||
raise KeyError(f"Path does not exist for given primary keys: {self._render_key_values(key_values)}")
|
||||
current[value] = {}
|
||||
next_node = current[value] # type: ignore
|
||||
if not isinstance(next_node, dict):
|
||||
raise ValueError(f"Internal structure corrupted: expected dict, got {type(next_node)!r}") # type: ignore
|
||||
current = next_node # type: ignore
|
||||
|
||||
# We should always return inside the loop.
|
||||
raise RuntimeError("Unreachable")
|
||||
|
||||
def _mutate_single(self, item: T, mode: MutationMode) -> None:
|
||||
"""Core mutation logic shared by insert, update, upsert, and delete."""
|
||||
self._ensure_item_type(item)
|
||||
key_values = self._extract_primary_key_values(item)
|
||||
|
||||
if mode in ("insert", "upsert"):
|
||||
parent, final_key = self._locate_node(key_values, create_missing=True)
|
||||
exists = final_key in parent
|
||||
|
||||
if mode == "insert":
|
||||
if exists:
|
||||
raise ValueError(f"Item already exists with primary key(s): {self._render_key_values(key_values)}")
|
||||
parent[final_key] = item
|
||||
self._size += 1
|
||||
else: # upsert
|
||||
if not exists:
|
||||
self._size += 1
|
||||
parent[final_key] = item
|
||||
|
||||
elif mode in ("update", "delete"):
|
||||
# For update/delete we must not create missing paths.
|
||||
try:
|
||||
parent, final_key = self._locate_node(key_values, create_missing=False)
|
||||
except KeyError:
|
||||
raise ValueError(
|
||||
f"Item does not exist with primary key(s): {self._render_key_values(key_values)}"
|
||||
) from None
|
||||
|
||||
if final_key not in parent:
|
||||
raise ValueError(f"Item does not exist with primary key(s): {self._render_key_values(key_values)}")
|
||||
|
||||
if mode == "update":
|
||||
parent[final_key] = item
|
||||
else: # delete
|
||||
del parent[final_key]
|
||||
self._size -= 1
|
||||
else:
|
||||
raise ValueError(f"Unknown mutation mode: {mode}")
|
||||
|
||||
def _iter_items(
|
||||
self,
|
||||
root: Optional[Mapping[Any, Any]] = None,
|
||||
filters: Optional[FilterMap] = None,
|
||||
must_filters: Optional[FilterMap] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
) -> Iterable[T]:
|
||||
"""Iterate over all items in the nested dictionary structure, optionally applying filters."""
|
||||
if root is None:
|
||||
root = self._items
|
||||
if not root:
|
||||
return
|
||||
stack: List[Mapping[Any, Any]] = [root]
|
||||
while stack:
|
||||
node = stack.pop()
|
||||
for value in node.values():
|
||||
# Leaf nodes contain items; intermediate nodes are dicts.
|
||||
if isinstance(value, self._item_type):
|
||||
if _item_matches_filters(value, filters, filter_logic, must_filters):
|
||||
yield value
|
||||
elif isinstance(value, dict):
|
||||
stack.append(value) # type: ignore
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Internal structure corrupted: expected dict or {self._item_type.__name__}, "
|
||||
f"got {type(value)!r}"
|
||||
)
|
||||
|
||||
def _iter_matching_items(
|
||||
self,
|
||||
filters: Optional[FilterMap],
|
||||
must_filters: Optional[FilterMap],
|
||||
filter_logic: Literal["and", "or"],
|
||||
) -> Iterable[T]:
|
||||
"""Efficiently iterate over items matching filters, using primary-key prefix when possible."""
|
||||
# Fast path: when optional filters can't form a prefix, fall back to scanning.
|
||||
if filter_logic != "and" and must_filters is None:
|
||||
return self._iter_items(filters=filters, must_filters=must_filters, filter_logic=filter_logic)
|
||||
|
||||
# Try to derive a primary-key prefix from exact filters.
|
||||
pk_values_prefix: List[Any] = []
|
||||
prefix_sources: List[FilterMap] = []
|
||||
if must_filters:
|
||||
prefix_sources.append(must_filters)
|
||||
if filter_logic == "and" and filters:
|
||||
prefix_sources.append(filters)
|
||||
|
||||
for pk in self._primary_keys:
|
||||
# combined_ops are: [{"exact": value}, {"within": [...]}, ...]
|
||||
combined_ops: List[FilterField] = []
|
||||
for source in prefix_sources:
|
||||
field_ops = source.get(pk) # type: ignore[union-attr]
|
||||
if field_ops:
|
||||
combined_ops.append(field_ops)
|
||||
if not combined_ops:
|
||||
break
|
||||
# Only allow a pure {"exact": value} constraint.
|
||||
exact_value: Any | None = None
|
||||
allow_prefix = True
|
||||
for ops in combined_ops:
|
||||
if set(ops.keys()) != {"exact"}:
|
||||
allow_prefix = False
|
||||
break
|
||||
candidate = ops.get("exact")
|
||||
if candidate is None:
|
||||
allow_prefix = False
|
||||
break
|
||||
if exact_value is not None and candidate != exact_value:
|
||||
# Contradictory exact filters mean no items can match.
|
||||
logger.warning(f"Contradictory exact filters for field '{pk}': {exact_value} != {candidate}")
|
||||
return ()
|
||||
exact_value = candidate
|
||||
|
||||
if not allow_prefix:
|
||||
break
|
||||
|
||||
value = exact_value
|
||||
if value is None:
|
||||
break
|
||||
pk_values_prefix.append(value)
|
||||
|
||||
if not pk_values_prefix:
|
||||
return self._iter_items(filters=filters, must_filters=must_filters, filter_logic=filter_logic)
|
||||
|
||||
try:
|
||||
if len(pk_values_prefix) == len(self._primary_keys):
|
||||
# All primary keys specified -> at most a single item.
|
||||
parent, final_key = self._locate_node(pk_values_prefix, create_missing=False)
|
||||
single_item = parent.get(final_key)
|
||||
if isinstance(single_item, self._item_type) and _item_matches_filters(
|
||||
single_item,
|
||||
filters,
|
||||
filter_logic,
|
||||
must_filters,
|
||||
):
|
||||
return (single_item,)
|
||||
return ()
|
||||
else:
|
||||
# Prefix of primary keys specified -> iterate only the subtree below that prefix.
|
||||
parent, final_key = self._locate_node(pk_values_prefix, create_missing=False)
|
||||
subtree = parent.get(final_key)
|
||||
if isinstance(subtree, dict):
|
||||
return self._iter_items(
|
||||
subtree, # type: ignore
|
||||
filters=filters,
|
||||
must_filters=must_filters,
|
||||
filter_logic=filter_logic,
|
||||
)
|
||||
return ()
|
||||
except KeyError:
|
||||
# No items exist for this primary-key prefix.
|
||||
return ()
|
||||
|
||||
async def query(
|
||||
self,
|
||||
filter: Optional[FilterOptions] = None,
|
||||
sort: Optional[SortOptions] = None,
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> PaginatedResult[T]:
|
||||
"""Query the collection with filters, sort order, and pagination.
|
||||
|
||||
Args:
|
||||
filter: Mapping of field name to operator dict along with the optional `_aggregate` logic.
|
||||
sort: Options describing which field to sort by and in which order.
|
||||
limit: Max number of items to return. Use -1 for "no limit".
|
||||
offset: Number of items to skip from the start of the *matching* items.
|
||||
"""
|
||||
filters, must_filters, filter_logic = _normalize_filter_options(filter)
|
||||
sort_by, sort_order = _resolve_sort_options(sort)
|
||||
items_iter: Iterable[T] = self._iter_matching_items(filters, must_filters, filter_logic)
|
||||
|
||||
# No sorting: stream through items and apply pagination on the fly.
|
||||
if not sort_by:
|
||||
matched_items: List[T] = []
|
||||
total_matched = 0
|
||||
|
||||
for item in items_iter:
|
||||
# Count every match for 'total'
|
||||
total_matched += 1
|
||||
|
||||
# Apply offset/limit window
|
||||
if total_matched <= offset:
|
||||
continue
|
||||
if limit != -1 and len(matched_items) >= limit:
|
||||
# Still need to finish iteration to get accurate total_matched.
|
||||
continue
|
||||
|
||||
matched_items.append(item)
|
||||
|
||||
return PaginatedResult(
|
||||
items=matched_items,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
total=total_matched,
|
||||
)
|
||||
|
||||
# With sorting: we must materialize all matching items to sort them.
|
||||
all_matches: List[T] = list(items_iter)
|
||||
|
||||
total_matched = len(all_matches)
|
||||
reverse = sort_order == "desc"
|
||||
all_matches.sort(key=lambda x: _get_sort_value(x, sort_by), reverse=reverse)
|
||||
|
||||
if limit == -1:
|
||||
paginated_items = all_matches[offset:]
|
||||
else:
|
||||
paginated_items = all_matches[offset : offset + limit]
|
||||
|
||||
return PaginatedResult(
|
||||
items=paginated_items,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
total=total_matched,
|
||||
)
|
||||
|
||||
async def get(
|
||||
self,
|
||||
filter: Optional[FilterOptions] = None,
|
||||
sort: Optional[SortOptions] = None,
|
||||
) -> Optional[T]:
|
||||
"""Return the first (or best-sorted) item that matches the given filters, or None."""
|
||||
filters, must_filters, filter_logic = _normalize_filter_options(filter)
|
||||
sort_by, sort_order = _resolve_sort_options(sort)
|
||||
items_iter: Iterable[T] = self._iter_matching_items(filters, must_filters, filter_logic)
|
||||
|
||||
if not sort_by:
|
||||
# Just return the first matching item, if any.
|
||||
for item in items_iter:
|
||||
return item
|
||||
return None
|
||||
|
||||
# Single-pass min/max according to sort_order.
|
||||
best_item: Optional[T] = None
|
||||
best_key: Any = None
|
||||
|
||||
for item in items_iter:
|
||||
key = _get_sort_value(item, sort_by)
|
||||
if best_item is None:
|
||||
best_item = item
|
||||
best_key = key
|
||||
continue
|
||||
|
||||
if sort_order == "asc":
|
||||
if key < best_key:
|
||||
best_item, best_key = item, key
|
||||
else:
|
||||
if key > best_key:
|
||||
best_item, best_key = item, key
|
||||
|
||||
return best_item
|
||||
|
||||
async def insert(self, items: Sequence[T]) -> None:
|
||||
"""Insert the given items.
|
||||
|
||||
Raises:
|
||||
ValueError: If any item with the same primary keys already exists.
|
||||
"""
|
||||
for item in items:
|
||||
self._mutate_single(item, mode="insert")
|
||||
|
||||
async def update(self, items: Sequence[T]) -> None:
|
||||
"""Update the given items.
|
||||
|
||||
Raises:
|
||||
ValueError: If any item with the given primary keys does not exist.
|
||||
"""
|
||||
for item in items:
|
||||
self._mutate_single(item, mode="update")
|
||||
|
||||
async def upsert(self, items: Sequence[T]) -> None:
|
||||
"""Upsert the given items (insert if missing, otherwise update)."""
|
||||
for item in items:
|
||||
self._mutate_single(item, mode="upsert")
|
||||
|
||||
async def delete(self, items: Sequence[T]) -> None:
|
||||
"""Delete the given items.
|
||||
|
||||
Raises:
|
||||
ValueError: If any item with the given primary keys does not exist.
|
||||
"""
|
||||
# We use a two-phase approach to avoid partial deletion if one fails:
|
||||
# first compute key_values to validate, then perform deletions.
|
||||
for item in items:
|
||||
# _mutate_single will validate existence and update size.
|
||||
self._mutate_single(item, mode="delete")
|
||||
|
||||
|
||||
class DequeBasedQueue(Queue[T]):
|
||||
"""Queue implementation backed by collections.deque.
|
||||
|
||||
Provides O(1) amortized enqueue (append) and dequeue (popleft).
|
||||
"""
|
||||
|
||||
def __init__(self, item_type: Type[T], items: Optional[Sequence[T]] = None):
|
||||
self._items: Deque[T] = deque()
|
||||
self._item_type: Type[T] = item_type
|
||||
if items:
|
||||
self._items.extend(items)
|
||||
|
||||
def item_type(self) -> Type[T]:
|
||||
return self._item_type
|
||||
|
||||
async def has(self, item: T) -> bool:
|
||||
if not isinstance(item, self._item_type):
|
||||
raise TypeError(f"Expected item of type {self._item_type.__name__}, got {type(item).__name__}")
|
||||
return item in self._items
|
||||
|
||||
async def enqueue(self, items: Sequence[T]) -> Sequence[T]:
|
||||
for item in items:
|
||||
if not isinstance(item, self._item_type):
|
||||
raise TypeError(f"Expected item of type {self._item_type.__name__}, got {type(item).__name__}")
|
||||
self._items.append(item)
|
||||
return items
|
||||
|
||||
async def dequeue(self, limit: int = 1) -> Sequence[T]:
|
||||
if limit <= 0:
|
||||
return []
|
||||
out: List[T] = []
|
||||
for _ in range(min(limit, len(self._items))):
|
||||
out.append(self._items.popleft())
|
||||
return out
|
||||
|
||||
async def peek(self, limit: int = 1) -> Sequence[T]:
|
||||
if limit <= 0:
|
||||
return []
|
||||
result: List[T] = []
|
||||
count = min(limit, len(self._items))
|
||||
for idx, item in enumerate(self._items):
|
||||
if idx >= count:
|
||||
break
|
||||
result.append(item)
|
||||
return result
|
||||
|
||||
def size(self) -> int:
|
||||
return len(self._items)
|
||||
|
||||
|
||||
class DictBasedKeyValue(KeyValue[K, V]):
|
||||
"""KeyValue implementation backed by a plain dictionary."""
|
||||
|
||||
def __init__(self, data: Optional[Mapping[K, V]] = None):
|
||||
self._values: Dict[K, V] = dict(data) if data else {}
|
||||
|
||||
async def has(self, key: K) -> bool:
|
||||
return key in self._values
|
||||
|
||||
async def get(self, key: K, default: V | None = None) -> V | None:
|
||||
return self._values.get(key, default)
|
||||
|
||||
async def set(self, key: K, value: V) -> None:
|
||||
self._values[key] = value
|
||||
|
||||
async def pop(self, key: K, default: V | None = None) -> V | None:
|
||||
return self._values.pop(key, default)
|
||||
|
||||
def size(self) -> int:
|
||||
return len(self._values)
|
||||
|
||||
|
||||
class InMemoryLightningCollections(LightningCollections):
|
||||
"""In-memory implementation of LightningCollections using Python data structures.
|
||||
|
||||
Serves as the storage base for [`InMemoryLightningStore`][agentlightning.InMemoryLightningStore].
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._lock = _LoopAwareAsyncLock()
|
||||
self._rollouts = ListBasedCollection(items=[], item_type=Rollout, primary_keys=["rollout_id"])
|
||||
self._attempts = ListBasedCollection(items=[], item_type=Attempt, primary_keys=["rollout_id", "attempt_id"])
|
||||
self._spans = ListBasedCollection(
|
||||
items=[], item_type=Span, primary_keys=["rollout_id", "attempt_id", "span_id"]
|
||||
)
|
||||
self._resources = ListBasedCollection(items=[], item_type=ResourcesUpdate, primary_keys=["resources_id"])
|
||||
self._workers = ListBasedCollection(items=[], item_type=Worker, primary_keys=["worker_id"])
|
||||
self._rollout_queue = DequeBasedQueue(items=[], item_type=str)
|
||||
self._span_sequence_ids = DictBasedKeyValue[str, int](data={}) # rollout_id -> sequence_id
|
||||
|
||||
@property
|
||||
def rollouts(self) -> ListBasedCollection[Rollout]:
|
||||
return self._rollouts
|
||||
|
||||
@property
|
||||
def attempts(self) -> ListBasedCollection[Attempt]:
|
||||
return self._attempts
|
||||
|
||||
@property
|
||||
def spans(self) -> ListBasedCollection[Span]:
|
||||
return self._spans
|
||||
|
||||
@property
|
||||
def resources(self) -> ListBasedCollection[ResourcesUpdate]:
|
||||
return self._resources
|
||||
|
||||
@property
|
||||
def workers(self) -> ListBasedCollection[Worker]:
|
||||
return self._workers
|
||||
|
||||
@property
|
||||
def rollout_queue(self) -> DequeBasedQueue[str]:
|
||||
return self._rollout_queue
|
||||
|
||||
@property
|
||||
def span_sequence_ids(self) -> DictBasedKeyValue[str, int]:
|
||||
return self._span_sequence_ids
|
||||
|
||||
@asynccontextmanager
|
||||
async def atomic(self, *args: Any, **kwargs: Any) -> AsyncGenerator[None, None]:
|
||||
async with self._lock:
|
||||
yield
|
||||
|
||||
async def evict_spans_for_rollout(self, rollout_id: str) -> None:
|
||||
"""Evict all spans for a given rollout ID.
|
||||
|
||||
Uses private API for efficiency.
|
||||
"""
|
||||
self._spans._items.pop(rollout_id, []) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
class _LoopAwareAsyncLock:
|
||||
"""Async lock that transparently rebinds to the current event loop.
|
||||
|
||||
The lock intentionally remains *thread-unsafe*: callers must only use it from
|
||||
one thread at a time. If multiple threads interact with the store, each
|
||||
thread gets its own event loop specific lock.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._locks: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock] = weakref.WeakKeyDictionary()
|
||||
|
||||
# When serializing and deserializing, we don't need to serialize the locks.
|
||||
# Because another process will have its own set of event loops and its own lock.
|
||||
def __getstate__(self) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def __setstate__(self, state: dict[str, Any]) -> None:
|
||||
self._locks = weakref.WeakKeyDictionary()
|
||||
|
||||
def _get_lock_for_current_loop(self) -> asyncio.Lock:
|
||||
loop = asyncio.get_running_loop()
|
||||
lock = self._locks.get(loop)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
self._locks[loop] = lock
|
||||
return lock
|
||||
|
||||
async def __aenter__(self) -> asyncio.Lock:
|
||||
lock = self._get_lock_for_current_loop()
|
||||
await lock.acquire()
|
||||
return lock
|
||||
|
||||
async def __aexit__(self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: Any) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
lock = self._locks.get(loop)
|
||||
if lock is None or not lock.locked():
|
||||
raise RuntimeError("Lock released without being acquired")
|
||||
lock.release()
|
||||
File diff suppressed because it is too large
Load Diff
+112
-755
@@ -3,15 +3,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import functools
|
||||
import hashlib
|
||||
import logging
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import weakref
|
||||
from collections import deque
|
||||
from collections.abc import Iterable
|
||||
from collections.abc import Mapping as MappingABC
|
||||
from typing import (
|
||||
@@ -23,76 +17,24 @@ from typing import (
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Set,
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from pydantic import BaseModel
|
||||
|
||||
from agentlightning.types import (
|
||||
Attempt,
|
||||
AttemptedRollout,
|
||||
AttemptStatus,
|
||||
NamedResources,
|
||||
ResourcesUpdate,
|
||||
Rollout,
|
||||
RolloutConfig,
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
)
|
||||
from agentlightning.types import AttemptedRollout, PaginatedResult, Rollout, Span
|
||||
|
||||
from .base import UNSET, LightningStore, Unset, is_finished, is_queuing
|
||||
from .utils import healthcheck, propagate_status
|
||||
from .base import LightningStoreCapabilities, is_finished, is_running
|
||||
from .collection import InMemoryLightningCollections
|
||||
from .collection_based import CollectionBasedLightningStore
|
||||
|
||||
T_callable = TypeVar("T_callable", bound=Callable[..., Any])
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _LoopAwareAsyncLock:
|
||||
"""Async lock that transparently rebinds to the current event loop.
|
||||
|
||||
The lock intentionally remains *thread-unsafe*: callers must only use it from
|
||||
one thread at a time. If multiple threads interact with the store, each
|
||||
thread gets its own event loop specific lock.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._locks: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock] = weakref.WeakKeyDictionary()
|
||||
|
||||
# When serializing and deserializing, we don't need to serialize the locks.
|
||||
# Because another process will have its own set of event loops and its own lock.
|
||||
def __getstate__(self) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def __setstate__(self, state: dict[str, Any]) -> None:
|
||||
self._locks = weakref.WeakKeyDictionary()
|
||||
|
||||
def _get_lock_for_current_loop(self) -> asyncio.Lock:
|
||||
loop = asyncio.get_running_loop()
|
||||
lock = self._locks.get(loop)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
self._locks[loop] = lock
|
||||
return lock
|
||||
|
||||
async def __aenter__(self) -> asyncio.Lock:
|
||||
lock = self._get_lock_for_current_loop()
|
||||
await lock.acquire()
|
||||
return lock
|
||||
|
||||
async def __aexit__(self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: Any) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
lock = self._locks.get(loop)
|
||||
if lock is None or not lock.locked():
|
||||
raise RuntimeError("Lock released without being acquired")
|
||||
lock.release()
|
||||
|
||||
|
||||
def estimate_model_size(obj: Any) -> int:
|
||||
"""Rough recursive size estimate for Pydantic BaseModel instances."""
|
||||
|
||||
@@ -108,53 +50,6 @@ def estimate_model_size(obj: Any) -> int:
|
||||
return sys.getsizeof(cast(object, obj))
|
||||
|
||||
|
||||
def _healthcheck_wrapper(func: T_callable) -> T_callable:
|
||||
"""
|
||||
Decorator to run the watchdog healthcheck **before** executing the decorated method.
|
||||
Only runs if the store has a watchdog configured.
|
||||
Prevents recursive healthcheck execution using a flag on the store instance.
|
||||
"""
|
||||
|
||||
@functools.wraps(func)
|
||||
async def wrapper(self: InMemoryLightningStore, *args: Any, **kwargs: Any) -> Any:
|
||||
# Check if healthcheck is already running to prevent recursion
|
||||
if getattr(self, "_healthcheck_running", False):
|
||||
# Skip healthcheck if already running
|
||||
return await func(self, *args, **kwargs)
|
||||
|
||||
# Set flag to prevent recursive healthcheck calls
|
||||
# This flag is not asyncio/thread-safe, but it doesn't matter
|
||||
self._healthcheck_running = True # type: ignore
|
||||
try:
|
||||
# The following methods should live inside one lock.
|
||||
await self._healthcheck() # pyright: ignore[reportPrivateUsage]
|
||||
finally:
|
||||
# Always clear the flag, even if healthcheck fails
|
||||
self._healthcheck_running = False # type: ignore
|
||||
|
||||
# Execute the original method
|
||||
# This should be outside the lock.
|
||||
return await func(self, *args, **kwargs)
|
||||
|
||||
return cast(T_callable, wrapper)
|
||||
|
||||
|
||||
def _generate_resources_id() -> str:
|
||||
short_id = hashlib.sha1(uuid.uuid4().bytes).hexdigest()[:12]
|
||||
return "rs-" + short_id
|
||||
|
||||
|
||||
def _generate_rollout_id() -> str:
|
||||
short_id = hashlib.sha1(uuid.uuid4().bytes).hexdigest()[:12]
|
||||
return "ro-" + short_id
|
||||
|
||||
|
||||
def _generate_attempt_id() -> str:
|
||||
"""We don't need that long because attempts are limited to rollouts."""
|
||||
short_id = hashlib.sha1(uuid.uuid4().bytes).hexdigest()[:8]
|
||||
return "at-" + short_id
|
||||
|
||||
|
||||
def _detect_total_memory_bytes() -> int:
|
||||
"""Best-effort detection of the total available system memory in bytes."""
|
||||
|
||||
@@ -168,7 +63,7 @@ def _detect_total_memory_bytes() -> int:
|
||||
return 8 * 1024**3
|
||||
|
||||
|
||||
class InMemoryLightningStore(LightningStore):
|
||||
class InMemoryLightningStore(CollectionBasedLightningStore[InMemoryLightningCollections]):
|
||||
"""
|
||||
In-memory implementation of LightningStore using Python data structures.
|
||||
Thread-safe and async-compatible but data is not persistent.
|
||||
@@ -192,19 +87,9 @@ class InMemoryLightningStore(LightningStore):
|
||||
safe_memory_threshold: float | int | None = None,
|
||||
span_size_estimator: Callable[[Span], int] | None = None,
|
||||
):
|
||||
self._lock = _LoopAwareAsyncLock()
|
||||
super().__init__(collections=InMemoryLightningCollections())
|
||||
|
||||
# Task queue and rollouts storage
|
||||
self._task_queue: deque[Rollout] = deque()
|
||||
self._rollouts: Dict[str, Rollout] = {}
|
||||
|
||||
# Resources storage (similar to legacy server.py)
|
||||
self._resources: Dict[str, ResourcesUpdate] = {}
|
||||
self._latest_resources_id: Optional[str] = None
|
||||
|
||||
# Spans storage
|
||||
self._spans: Dict[str, List[Span]] = {} # rollout_id -> list of spans
|
||||
self._span_sequence_ids: Dict[str, int] = Counter() # rollout_id -> sequence_id
|
||||
self._start_time_by_rollout: Dict[str, float] = {}
|
||||
self._span_bytes_by_rollout: Dict[str, int] = Counter()
|
||||
self._total_span_bytes: int = 0
|
||||
self._evicted_rollout_span_sets: Set[str] = set()
|
||||
@@ -236,368 +121,111 @@ class InMemoryLightningStore(LightningStore):
|
||||
raise ValueError("safe_memory_threshold must be smaller than eviction_memory_threshold")
|
||||
self._custom_span_size_estimator = span_size_estimator
|
||||
|
||||
# Attempt tracking
|
||||
self._attempts: Dict[str, List[Attempt]] = {} # rollout_id -> list of attempts
|
||||
|
||||
# Completion tracking for wait_for_rollouts (cross-loop safe)
|
||||
self._completion_events: Dict[str, threading.Event] = {}
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def start_rollout(
|
||||
self,
|
||||
input: TaskInput,
|
||||
mode: Literal["train", "val", "test"] | None = None,
|
||||
resources_id: str | None = None,
|
||||
config: RolloutConfig | None = None,
|
||||
metadata: Dict[str, Any] | None = None,
|
||||
) -> AttemptedRollout:
|
||||
"""Notify the store that I'm about to run a rollout.
|
||||
# Running rollouts cache, including preparing and running rollouts
|
||||
self._running_rollout_ids: Set[str] = set()
|
||||
|
||||
See [`LightningStore.start_rollout()`][agentlightning.LightningStore.start_rollout] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
rollout_id = _generate_rollout_id()
|
||||
current_time = time.time()
|
||||
@property
|
||||
def capabilities(self) -> LightningStoreCapabilities:
|
||||
"""Return the capabilities of the store."""
|
||||
return LightningStoreCapabilities(
|
||||
thread_safe=False,
|
||||
async_safe=True,
|
||||
zero_copy=False,
|
||||
otlp_traces=False,
|
||||
)
|
||||
|
||||
rollout_config = config.model_copy(deep=True) if config is not None else RolloutConfig()
|
||||
rollout_metadata = dict(metadata) if metadata is not None else {}
|
||||
async def wait_for_rollout(self, rollout_id: str, timeout: Optional[float] = None) -> Optional[Rollout]:
|
||||
"""Wait for a specific rollout to complete with a timeout."""
|
||||
async with self.collections.atomic():
|
||||
rollout = await self.collections.rollouts.get({"rollout_id": {"exact": rollout_id}})
|
||||
if rollout and is_finished(rollout):
|
||||
return rollout
|
||||
|
||||
rollout = Rollout(
|
||||
rollout_id=rollout_id,
|
||||
input=input,
|
||||
mode=mode,
|
||||
resources_id=resources_id or self._latest_resources_id,
|
||||
start_time=current_time,
|
||||
status="preparing",
|
||||
config=rollout_config,
|
||||
metadata=rollout_metadata,
|
||||
)
|
||||
|
||||
# Create the initial attempt
|
||||
attempt_id = _generate_attempt_id()
|
||||
attempt = Attempt(
|
||||
rollout_id=rollout.rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
sequence_id=1,
|
||||
start_time=current_time,
|
||||
status="preparing",
|
||||
)
|
||||
|
||||
self._attempts[rollout.rollout_id] = [attempt]
|
||||
self._rollouts[rollout.rollout_id] = rollout
|
||||
|
||||
# Manully added rollout is not added to task queue. It's already preparing
|
||||
self._completion_events.setdefault(rollout.rollout_id, threading.Event())
|
||||
|
||||
return AttemptedRollout(**rollout.model_dump(), attempt=attempt)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def enqueue_rollout(
|
||||
self,
|
||||
input: TaskInput,
|
||||
mode: Literal["train", "val", "test"] | None = None,
|
||||
resources_id: str | None = None,
|
||||
config: RolloutConfig | None = None,
|
||||
metadata: Dict[str, Any] | None = None,
|
||||
) -> Rollout:
|
||||
"""Adds a new task to the queue with specific metadata and returns the rollout.
|
||||
|
||||
See [`LightningStore.enqueue_rollout()`][agentlightning.LightningStore.enqueue_rollout] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
rollout_id = _generate_rollout_id()
|
||||
current_time = time.time()
|
||||
|
||||
rollout_config = config.model_copy(deep=True) if config is not None else RolloutConfig()
|
||||
rollout_metadata = dict(metadata) if metadata is not None else {}
|
||||
|
||||
rollout = Rollout(
|
||||
rollout_id=rollout_id,
|
||||
input=input,
|
||||
mode=mode,
|
||||
resources_id=resources_id or self._latest_resources_id,
|
||||
start_time=current_time,
|
||||
status="queuing", # should be queuing
|
||||
config=rollout_config,
|
||||
metadata=rollout_metadata,
|
||||
)
|
||||
|
||||
self._rollouts[rollout.rollout_id] = rollout
|
||||
self._task_queue.append(rollout) # add it to the end of the queue
|
||||
self._completion_events.setdefault(rollout.rollout_id, threading.Event())
|
||||
|
||||
return rollout
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
"""Retrieves the next task from the queue without blocking.
|
||||
Returns `None` if the queue is empty.
|
||||
|
||||
Will set the rollout status to preparing and create a new attempt.
|
||||
|
||||
See [`LightningStore.dequeue_rollout()`][agentlightning.LightningStore.dequeue_rollout] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
# Keep looking until we find a rollout that's still in queuing status
|
||||
# or the queue is empty
|
||||
while self._task_queue:
|
||||
rollout = self._task_queue.popleft()
|
||||
|
||||
# Check if rollout is still in a queuing state
|
||||
# (it might have been updated to a different status while in queue)
|
||||
if is_queuing(rollout):
|
||||
# Update status to preparing
|
||||
rollout.status = "preparing"
|
||||
|
||||
# Create a new attempt (could be first attempt or retry)
|
||||
attempt_id = _generate_attempt_id()
|
||||
current_time = time.time()
|
||||
|
||||
# Get existing attempts to determine sequence number
|
||||
existing_attempts = self._attempts.get(rollout.rollout_id, [])
|
||||
sequence_id = len(existing_attempts) + 1
|
||||
|
||||
attempt = Attempt(
|
||||
rollout_id=rollout.rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
sequence_id=sequence_id,
|
||||
start_time=current_time,
|
||||
status="preparing",
|
||||
)
|
||||
|
||||
if rollout.rollout_id not in self._attempts:
|
||||
self._attempts[rollout.rollout_id] = []
|
||||
self._attempts[rollout.rollout_id].append(attempt)
|
||||
|
||||
return AttemptedRollout(**rollout.model_dump(), attempt=attempt)
|
||||
|
||||
# If not in queuing state, skip this rollout and continue
|
||||
# (it was updated externally and should not be processed)
|
||||
|
||||
# No valid rollouts found
|
||||
if timeout is not None and timeout <= 0:
|
||||
return None
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
|
||||
"""Creates a new attempt for a given rollout ID and return the attempt details.
|
||||
# If not completed and we have an event, wait for completion
|
||||
if rollout_id in self._completion_events:
|
||||
evt = self._completion_events[rollout_id]
|
||||
|
||||
See [`LightningStore.start_attempt()`][agentlightning.LightningStore.start_attempt] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
# Get the rollout
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if not rollout:
|
||||
raise ValueError(f"Rollout {rollout_id} not found")
|
||||
|
||||
# Get existing attempts to determine sequence number
|
||||
existing_attempts = self._attempts.get(rollout_id, [])
|
||||
sequence_id = len(existing_attempts) + 1
|
||||
|
||||
# We don't care whether the max attempts have reached or not
|
||||
# This attempt is from user trigger
|
||||
|
||||
# Create new attempt
|
||||
attempt_id = _generate_attempt_id()
|
||||
current_time = time.time()
|
||||
|
||||
attempt = Attempt(
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
sequence_id=sequence_id,
|
||||
start_time=current_time,
|
||||
status="preparing",
|
||||
)
|
||||
|
||||
# Add attempt to storage
|
||||
if rollout_id not in self._attempts:
|
||||
self._attempts[rollout_id] = []
|
||||
self._attempts[rollout_id].append(attempt)
|
||||
|
||||
self._completion_events.setdefault(rollout.rollout_id, threading.Event())
|
||||
|
||||
return AttemptedRollout(**rollout.model_dump(), attempt=attempt)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def query_rollouts(
|
||||
self, *, status: Optional[Sequence[RolloutStatus]] = None, rollout_ids: Optional[Sequence[str]] = None
|
||||
) -> List[Rollout]:
|
||||
"""Retrieves rollouts filtered by their status and rollout ids.
|
||||
If no status is provided, returns all rollouts.
|
||||
|
||||
See [`LightningStore.query_rollouts()`][agentlightning.LightningStore.query_rollouts] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
rollouts = list(self._rollouts.values())
|
||||
|
||||
# Filter by rollout_ids if provided
|
||||
if rollout_ids is not None:
|
||||
rollout_ids_set = set(rollout_ids)
|
||||
rollouts = [rollout for rollout in rollouts if rollout.rollout_id in rollout_ids_set]
|
||||
|
||||
# Filter by status if provided
|
||||
if status is not None:
|
||||
status_set = set(status)
|
||||
rollouts = [rollout for rollout in rollouts if rollout.status in status_set]
|
||||
|
||||
return rollouts
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def get_rollout_by_id(self, rollout_id: str) -> Optional[Rollout]:
|
||||
"""Retrieves a specific rollout by its ID.
|
||||
|
||||
See [`LightningStore.get_rollout_by_id()`][agentlightning.LightningStore.get_rollout_by_id] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
return self._rollouts.get(rollout_id)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def query_attempts(self, rollout_id: str) -> List[Attempt]:
|
||||
"""Retrieves all attempts associated with a specific rollout ID.
|
||||
Returns an empty list if no attempts are found.
|
||||
|
||||
See [`LightningStore.query_attempts()`][agentlightning.LightningStore.query_attempts] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
return self._attempts.get(rollout_id, [])
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def get_latest_attempt(self, rollout_id: str) -> Optional[Attempt]:
|
||||
"""Retrieves the latest attempt for a given rollout ID.
|
||||
|
||||
See [`LightningStore.get_latest_attempt()`][agentlightning.LightningStore.get_latest_attempt] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
attempts = self._attempts.get(rollout_id, [])
|
||||
if not attempts:
|
||||
return None
|
||||
return max(attempts, key=lambda a: a.sequence_id)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def add_resources(self, resources: NamedResources) -> ResourcesUpdate:
|
||||
"""Stores a new version of named resources and sets it as the latest.
|
||||
|
||||
See [`LightningStore.add_resources()`][agentlightning.LightningStore.add_resources] for semantics.
|
||||
"""
|
||||
resources_id = _generate_resources_id()
|
||||
async with self._lock:
|
||||
update = ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
self._resources[resources_id] = update
|
||||
self._latest_resources_id = resources_id
|
||||
return update
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def update_resources(self, resources_id: str, resources: NamedResources) -> ResourcesUpdate:
|
||||
"""
|
||||
Safely stores a new version of named resources and sets it as the latest.
|
||||
|
||||
See [`LightningStore.update_resources()`][agentlightning.LightningStore.update_resources] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
update = ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
self._resources[resources_id] = update
|
||||
self._latest_resources_id = resources_id
|
||||
return update
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def get_resources_by_id(self, resources_id: str) -> Optional[ResourcesUpdate]:
|
||||
"""Retrieves a specific version of named resources by its ID.
|
||||
|
||||
See [`LightningStore.get_resources_by_id()`][agentlightning.LightningStore.get_resources_by_id] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
return self._resources.get(resources_id)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def get_latest_resources(self) -> Optional[ResourcesUpdate]:
|
||||
"""Retrieves the latest version of named resources.
|
||||
|
||||
See [`LightningStore.get_latest_resources()`][agentlightning.LightningStore.get_latest_resources] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
if self._latest_resources_id:
|
||||
return self._resources.get(self._latest_resources_id)
|
||||
return None
|
||||
|
||||
async def get_next_span_sequence_id(self, rollout_id: str, attempt_id: str) -> int:
|
||||
"""Get the next span sequence ID for a given rollout and attempt.
|
||||
The number is strictly increasing for each rollout.
|
||||
The store will not issue the same sequence ID twice.
|
||||
|
||||
See [`LightningStore.get_next_span_sequence_id()`][agentlightning.LightningStore.get_next_span_sequence_id] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
self._span_sequence_ids[rollout_id] += 1
|
||||
return self._span_sequence_ids[rollout_id]
|
||||
|
||||
async def add_span(self, span: Span) -> Span:
|
||||
"""Persist a pre-converted span.
|
||||
|
||||
See [`LightningStore.add_span()`][agentlightning.LightningStore.add_span] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
self._span_sequence_ids[span.rollout_id] = max(self._span_sequence_ids[span.rollout_id], span.sequence_id)
|
||||
return await self._add_span_unlocked(span)
|
||||
|
||||
async def add_otel_span(
|
||||
self, rollout_id: str, attempt_id: str, readable_span: ReadableSpan, sequence_id: int | None = None
|
||||
) -> Span:
|
||||
"""Add an opentelemetry span to the store.
|
||||
|
||||
See [`LightningStore.add_otel_span()`][agentlightning.LightningStore.add_otel_span] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
if sequence_id is None:
|
||||
# Issue a new sequence ID for the rollout
|
||||
self._span_sequence_ids[rollout_id] += 1
|
||||
sequence_id = self._span_sequence_ids[rollout_id]
|
||||
# Wait for the event with proper timeout handling
|
||||
# evt.wait() returns True if event was set, False if timeout occurred
|
||||
if timeout is None:
|
||||
# Wait indefinitely by polling with finite timeouts
|
||||
# This allows threads to exit cleanly on shutdown
|
||||
while True:
|
||||
result = await asyncio.to_thread(evt.wait, 10.0) # Poll every 10 seconds
|
||||
if result: # Event was set
|
||||
break
|
||||
# Loop and check again (continues indefinitely since timeout=None)
|
||||
else:
|
||||
# Comes from a provided sequence ID
|
||||
# Make sure our counter is strictly increasing
|
||||
self._span_sequence_ids[rollout_id] = max(self._span_sequence_ids[rollout_id], sequence_id)
|
||||
# Wait with the specified timeout
|
||||
result = await asyncio.to_thread(evt.wait, timeout)
|
||||
|
||||
span = Span.from_opentelemetry(
|
||||
readable_span, rollout_id=rollout_id, attempt_id=attempt_id, sequence_id=sequence_id
|
||||
# If event was set (not timeout), check if rollout is finished
|
||||
if result:
|
||||
async with self.collections.atomic():
|
||||
rollout = await self.collections.rollouts.get({"rollout_id": {"exact": rollout_id}})
|
||||
if rollout and is_finished(rollout):
|
||||
return rollout
|
||||
|
||||
return None
|
||||
|
||||
async def on_rollout_update(self, rollout: Rollout) -> None:
|
||||
"""Update the running rollout ids set when the rollout updates."""
|
||||
if is_running(rollout):
|
||||
self._running_rollout_ids.add(rollout.rollout_id)
|
||||
else:
|
||||
self._running_rollout_ids.discard(rollout.rollout_id)
|
||||
|
||||
if is_finished(rollout):
|
||||
self._completion_events.setdefault(rollout.rollout_id, threading.Event())
|
||||
self._completion_events[rollout.rollout_id].set()
|
||||
else:
|
||||
self._completion_events.setdefault(rollout.rollout_id, threading.Event())
|
||||
# Rollout status can never transition from finished to running (unlike attempt)
|
||||
# so we don't need to clear the completion event even in case of retrying.
|
||||
|
||||
if rollout.rollout_id not in self._start_time_by_rollout:
|
||||
self._start_time_by_rollout[rollout.rollout_id] = rollout.start_time
|
||||
|
||||
async def get_running_rollouts(self) -> List[AttemptedRollout]:
|
||||
"""Accelerated version of `get_running_rollouts` for in-memory store. Used for healthcheck."""
|
||||
rollouts = await self.collections.rollouts.query(
|
||||
filter={"rollout_id": {"within": list(self._running_rollout_ids)}}
|
||||
)
|
||||
running_rollouts: List[AttemptedRollout] = []
|
||||
for rollout in rollouts.items:
|
||||
latest_attempt = await self.collections.attempts.get(
|
||||
filter={"rollout_id": {"exact": rollout.rollout_id}},
|
||||
sort={"name": "sequence_id", "order": "desc"},
|
||||
)
|
||||
await self._add_span_unlocked(span)
|
||||
return span
|
||||
if not latest_attempt:
|
||||
# The rollout is running but has no attempts, this should not happen
|
||||
logger.error(f"Rollout {rollout.rollout_id} is running but has no attempts")
|
||||
continue
|
||||
running_rollouts.append(AttemptedRollout(**rollout.model_dump(), attempt=latest_attempt))
|
||||
return running_rollouts
|
||||
|
||||
async def query_spans(
|
||||
self,
|
||||
rollout_id: str,
|
||||
attempt_id: str | Literal["latest"] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> PaginatedResult[Span]:
|
||||
if rollout_id in self._evicted_rollout_span_sets:
|
||||
raise RuntimeError(f"Spans for rollout {rollout_id} have been evicted")
|
||||
return await super().query_spans(rollout_id, attempt_id, **kwargs)
|
||||
|
||||
async def _add_span_unlocked(self, span: Span) -> Span:
|
||||
rollout = self._rollouts.get(span.rollout_id)
|
||||
if not rollout:
|
||||
raise ValueError(f"Rollout {span.rollout_id} not found")
|
||||
attempts = self._attempts.get(span.rollout_id, [])
|
||||
current_attempt = next((a for a in attempts if a.attempt_id == span.attempt_id), None)
|
||||
latest_attempt = max(attempts, key=lambda a: a.sequence_id) if attempts else None
|
||||
if not current_attempt:
|
||||
raise ValueError(f"Attempt {span.attempt_id} not found for rollout {span.rollout_id}")
|
||||
if not latest_attempt:
|
||||
raise ValueError(f"No attempts found for rollout {span.rollout_id}")
|
||||
"""In-memory store needs to maintain the span data in memory, and evict spans when memory is low."""
|
||||
|
||||
if span.rollout_id not in self._spans:
|
||||
self._spans[span.rollout_id] = []
|
||||
self._spans[span.rollout_id].append(span)
|
||||
await super()._add_span_unlocked(span)
|
||||
self._account_span_size(span)
|
||||
self._maybe_evict_spans()
|
||||
|
||||
# Update attempt heartbeat
|
||||
current_attempt.last_heartbeat_time = time.time()
|
||||
if current_attempt.status in ["preparing", "unresponsive"]:
|
||||
current_attempt.status = "running"
|
||||
|
||||
# If the status has already timed out or failed, do not change it
|
||||
|
||||
# Update rollout status if it's the latest attempt
|
||||
if current_attempt == latest_attempt:
|
||||
if rollout.status == "preparing":
|
||||
rollout.status = "running"
|
||||
elif rollout.status in ["queuing", "requeuing"]:
|
||||
try:
|
||||
self._task_queue.remove(rollout)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
f"Trying to remove rollout {rollout.rollout_id} from the queue but it's not in the queue."
|
||||
)
|
||||
rollout.status = "running"
|
||||
await self._maybe_evict_spans()
|
||||
|
||||
return span
|
||||
|
||||
@@ -641,19 +269,18 @@ class InMemoryLightningStore(LightningStore):
|
||||
self._total_span_bytes += size
|
||||
return size
|
||||
|
||||
def _maybe_evict_spans(self) -> None:
|
||||
async def _maybe_evict_spans(self) -> None:
|
||||
if self._total_span_bytes <= self._eviction_threshold_bytes:
|
||||
return
|
||||
|
||||
candidates: List[tuple[float, str]] = []
|
||||
for rollout_id, spans in self._spans.items():
|
||||
if not spans:
|
||||
continue
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
start_time = rollout.start_time if rollout is not None else (spans[0].start_time or 0.0)
|
||||
candidates.append((start_time, rollout_id))
|
||||
|
||||
candidates.sort(key=lambda item: item[0])
|
||||
logger.info(
|
||||
f"Total span bytes: {self._total_span_bytes}, eviction threshold: {self._eviction_threshold_bytes}, "
|
||||
f"safe threshold: {self._safe_threshold_bytes}. Evicting spans..."
|
||||
)
|
||||
candidates: List[tuple[float, str]] = [
|
||||
(start_time, rollout_id) for rollout_id, start_time in self._start_time_by_rollout.items()
|
||||
]
|
||||
candidates.sort()
|
||||
|
||||
logger.info(f"Evicting spans for {len(candidates)} rollouts to free up memory...")
|
||||
memory_consumed_before = self._total_span_bytes
|
||||
@@ -661,283 +288,13 @@ class InMemoryLightningStore(LightningStore):
|
||||
if self._total_span_bytes <= self._safe_threshold_bytes:
|
||||
break
|
||||
logger.debug(f"Evicting spans for rollout {rollout_id} to free up memory...")
|
||||
self._evict_spans_for_rollout(rollout_id)
|
||||
await self._evict_spans_for_rollout(rollout_id)
|
||||
logger.info(f"Freed up {memory_consumed_before - self._total_span_bytes} bytes of memory")
|
||||
|
||||
def _evict_spans_for_rollout(self, rollout_id: str) -> None:
|
||||
spans = self._spans.pop(rollout_id, [])
|
||||
if not spans:
|
||||
return
|
||||
async def _evict_spans_for_rollout(self, rollout_id: str) -> None:
|
||||
await self.collections.evict_spans_for_rollout(rollout_id)
|
||||
removed_bytes = self._span_bytes_by_rollout.pop(rollout_id, 0)
|
||||
self._total_span_bytes = max(self._total_span_bytes - removed_bytes, 0)
|
||||
self._evicted_rollout_span_sets.add(rollout_id)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def wait_for_rollouts(self, *, rollout_ids: List[str], timeout: Optional[float] = None) -> List[Rollout]:
|
||||
"""Wait for specified rollouts to complete with a timeout.
|
||||
Returns the completed rollouts, potentially incomplete if timeout is reached.
|
||||
|
||||
This method does not change the state of the store.
|
||||
|
||||
See [`LightningStore.wait_for_rollouts()`][agentlightning.LightningStore.wait_for_rollouts] for semantics.
|
||||
"""
|
||||
completed_rollouts: List[Rollout] = []
|
||||
|
||||
async def wait_for_rollout(rollout_id: str):
|
||||
# First check if already completed
|
||||
async with self._lock:
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if rollout and is_finished(rollout):
|
||||
completed_rollouts.append(rollout)
|
||||
return
|
||||
|
||||
# No timeout, return immediately
|
||||
if timeout is not None and timeout <= 0:
|
||||
return
|
||||
|
||||
# If not completed and we have an event, wait for completion
|
||||
if rollout_id in self._completion_events:
|
||||
evt = self._completion_events[rollout_id]
|
||||
|
||||
# Wait for the event with proper timeout handling
|
||||
# evt.wait() returns True if event was set, False if timeout occurred
|
||||
if timeout is None:
|
||||
# Wait indefinitely by polling with finite timeouts
|
||||
# This allows threads to exit cleanly on shutdown
|
||||
while True:
|
||||
result = await asyncio.to_thread(evt.wait, 10.0) # Poll every 10 seconds
|
||||
if result: # Event was set
|
||||
break
|
||||
# Loop and check again (continues indefinitely since timeout=None)
|
||||
else:
|
||||
# Wait with the specified timeout
|
||||
result = await asyncio.to_thread(evt.wait, timeout)
|
||||
|
||||
# If event was set (not timeout), check if rollout is finished
|
||||
if result:
|
||||
async with self._lock:
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if rollout and is_finished(rollout):
|
||||
completed_rollouts.append(rollout)
|
||||
|
||||
# Rollout not found, return
|
||||
|
||||
# Wait for all rollouts concurrently
|
||||
await asyncio.gather(*[wait_for_rollout(rid) for rid in rollout_ids], return_exceptions=True)
|
||||
|
||||
return completed_rollouts
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def query_spans(self, rollout_id: str, attempt_id: str | Literal["latest"] | None = None) -> List[Span]:
|
||||
"""
|
||||
Query and retrieve all spans associated with a specific rollout ID.
|
||||
Returns an empty list if no spans are found.
|
||||
|
||||
See [`LightningStore.query_spans()`][agentlightning.LightningStore.query_spans] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
if rollout_id in self._evicted_rollout_span_sets:
|
||||
raise RuntimeError(f"Spans for rollout {rollout_id} have been evicted")
|
||||
spans = self._spans.get(rollout_id, [])
|
||||
if attempt_id is None:
|
||||
return spans
|
||||
elif attempt_id == "latest":
|
||||
# Find the latest attempt_id
|
||||
if not spans:
|
||||
return []
|
||||
latest_attempt = max(spans, key=lambda s: s.sequence_id if s.attempt_id else "").attempt_id
|
||||
return [s for s in spans if s.attempt_id == latest_attempt]
|
||||
else:
|
||||
return [s for s in spans if s.attempt_id == attempt_id]
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def update_rollout(
|
||||
self,
|
||||
rollout_id: str,
|
||||
input: TaskInput | Unset = UNSET,
|
||||
mode: Optional[Literal["train", "val", "test"]] | Unset = UNSET,
|
||||
resources_id: Optional[str] | Unset = UNSET,
|
||||
status: RolloutStatus | Unset = UNSET,
|
||||
config: RolloutConfig | Unset = UNSET,
|
||||
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
|
||||
) -> Rollout:
|
||||
"""Update the rollout status and related metadata.
|
||||
|
||||
See [`LightningStore.update_rollout()`][agentlightning.LightningStore.update_rollout] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
return await self._update_rollout_unlocked(
|
||||
rollout_id=rollout_id,
|
||||
input=input,
|
||||
mode=mode,
|
||||
resources_id=resources_id,
|
||||
status=status,
|
||||
config=config,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
@_healthcheck_wrapper
|
||||
async def update_attempt(
|
||||
self,
|
||||
rollout_id: str,
|
||||
attempt_id: str | Literal["latest"],
|
||||
status: AttemptStatus | Unset = UNSET,
|
||||
worker_id: str | Unset = UNSET,
|
||||
last_heartbeat_time: float | Unset = UNSET,
|
||||
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
|
||||
) -> Attempt:
|
||||
"""Update a specific or latest attempt for a given rollout.
|
||||
|
||||
See [`LightningStore.update_attempt()`][agentlightning.LightningStore.update_attempt] for semantics.
|
||||
"""
|
||||
async with self._lock:
|
||||
attempt = await self._update_attempt_unlocked(
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
status=status,
|
||||
worker_id=worker_id,
|
||||
last_heartbeat_time=last_heartbeat_time,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
return attempt
|
||||
|
||||
async def _update_rollout_unlocked(
|
||||
self,
|
||||
rollout_id: str,
|
||||
input: TaskInput | Unset = UNSET,
|
||||
mode: Optional[Literal["train", "val", "test"]] | Unset = UNSET,
|
||||
resources_id: Optional[str] | Unset = UNSET,
|
||||
status: RolloutStatus | Unset = UNSET,
|
||||
config: RolloutConfig | Unset = UNSET,
|
||||
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
|
||||
) -> Rollout:
|
||||
# No lock inside this one.
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if not rollout:
|
||||
raise ValueError(f"Rollout {rollout_id} not found")
|
||||
|
||||
# Update fields if they are not UNSET
|
||||
if not isinstance(input, Unset):
|
||||
rollout.input = input
|
||||
if not isinstance(mode, Unset):
|
||||
rollout.mode = mode
|
||||
if not isinstance(resources_id, Unset):
|
||||
rollout.resources_id = resources_id
|
||||
if not isinstance(status, Unset):
|
||||
rollout.status = status
|
||||
if not isinstance(config, Unset):
|
||||
rollout.config = config
|
||||
if not isinstance(metadata, Unset):
|
||||
rollout.metadata = metadata
|
||||
|
||||
# Set end time for finished rollouts
|
||||
# Rollout is only finished when it succeeded or fail with no more retries.
|
||||
if not isinstance(status, Unset) and is_finished(rollout):
|
||||
rollout.end_time = time.time()
|
||||
# Signal completion
|
||||
if rollout_id in self._completion_events:
|
||||
self._completion_events[rollout_id].set()
|
||||
|
||||
# If requeuing, add back to queue
|
||||
elif is_queuing(rollout) and rollout not in self._task_queue:
|
||||
self._task_queue.append(rollout)
|
||||
|
||||
# If the rollout is no longer in a queueing state, remove it from the queue.
|
||||
if not isinstance(status, Unset) and not is_queuing(rollout) and rollout in self._task_queue:
|
||||
try:
|
||||
self._task_queue.remove(rollout)
|
||||
except ValueError:
|
||||
# Another coroutine may have already removed the rollout from the queue.
|
||||
logger.warning(
|
||||
f"Trying to remove rollout {rollout.rollout_id} from the queue but it's not in the queue."
|
||||
)
|
||||
|
||||
# Re-validate the rollout to ensure legality
|
||||
Rollout.model_validate(rollout.model_dump())
|
||||
|
||||
return rollout
|
||||
|
||||
async def _update_attempt_unlocked(
|
||||
self,
|
||||
rollout_id: str,
|
||||
attempt_id: str | Literal["latest"],
|
||||
status: AttemptStatus | Unset = UNSET,
|
||||
worker_id: str | Unset = UNSET,
|
||||
last_heartbeat_time: float | Unset = UNSET,
|
||||
metadata: Optional[Dict[str, Any]] | Unset = UNSET,
|
||||
) -> Attempt:
|
||||
# No lock, but with status propagation.
|
||||
rollout = self._rollouts.get(rollout_id)
|
||||
if not rollout:
|
||||
raise ValueError(f"Rollout {rollout_id} not found")
|
||||
|
||||
attempts = self._attempts.get(rollout_id, [])
|
||||
if not attempts:
|
||||
raise ValueError(f"No attempts found for rollout {rollout_id}")
|
||||
|
||||
latest_attempt = max(attempts, key=lambda a: a.sequence_id)
|
||||
|
||||
# Find the attempt to update
|
||||
if attempt_id == "latest":
|
||||
attempt = latest_attempt
|
||||
else:
|
||||
attempt = next((a for a in attempts if a.attempt_id == attempt_id), None)
|
||||
if not attempt:
|
||||
raise ValueError(f"Attempt {attempt_id} not found for rollout {rollout_id}")
|
||||
|
||||
# Update fields if they are not UNSET
|
||||
if not isinstance(status, Unset):
|
||||
attempt.status = status
|
||||
# Also update end_time if the status indicates completion
|
||||
if status in ["failed", "succeeded"]:
|
||||
attempt.end_time = time.time()
|
||||
if not isinstance(worker_id, Unset):
|
||||
attempt.worker_id = worker_id
|
||||
if not isinstance(last_heartbeat_time, Unset):
|
||||
attempt.last_heartbeat_time = last_heartbeat_time
|
||||
if not isinstance(metadata, Unset):
|
||||
attempt.metadata = metadata
|
||||
|
||||
# Re-validate the attempt to ensure legality
|
||||
Attempt.model_validate(attempt.model_dump())
|
||||
|
||||
if attempt == latest_attempt:
|
||||
|
||||
async def _update_status(rollout_id: str, status: RolloutStatus) -> Rollout:
|
||||
return await self._update_rollout_unlocked(rollout_id, status=status)
|
||||
|
||||
# Propagate the status to the rollout
|
||||
await propagate_status(
|
||||
_update_status,
|
||||
attempt,
|
||||
rollout.config,
|
||||
)
|
||||
|
||||
return attempt
|
||||
|
||||
async def _healthcheck(self) -> None:
|
||||
"""Perform healthcheck against all running rollouts in the store."""
|
||||
async with self._lock:
|
||||
running_rollouts: List[AttemptedRollout] = []
|
||||
for rollout in self._rollouts.values():
|
||||
if rollout.status in ["preparing", "running"]:
|
||||
all_attempts = self._attempts.get(rollout.rollout_id, [])
|
||||
if not all_attempts:
|
||||
# The rollout is running but has no attempts, this should not happen
|
||||
logger.error(f"Rollout {rollout.rollout_id} is running but has no attempts")
|
||||
continue
|
||||
latest_attempt = max(all_attempts, key=lambda a: a.sequence_id)
|
||||
running_rollouts.append(AttemptedRollout(**rollout.model_dump(), attempt=latest_attempt))
|
||||
|
||||
async def _update_attempt_status(rollout_id: str, attempt_id: str, status: AttemptStatus) -> Attempt:
|
||||
return await self._update_attempt_unlocked(rollout_id, attempt_id, status=status)
|
||||
|
||||
async def _update_rollout_status(rollout_id: str, status: RolloutStatus) -> Rollout:
|
||||
return await self._update_rollout_unlocked(rollout_id, status=status)
|
||||
|
||||
await healthcheck(
|
||||
running_rollouts,
|
||||
_update_rollout_status,
|
||||
_update_attempt_status,
|
||||
)
|
||||
if removed_bytes > 0:
|
||||
# There is something removed for real
|
||||
self._total_span_bytes = max(self._total_span_bytes - removed_bytes, 0)
|
||||
self._evicted_rollout_span_sets.add(rollout_id)
|
||||
|
||||
@@ -18,9 +18,11 @@ from agentlightning.types import (
|
||||
RolloutStatus,
|
||||
Span,
|
||||
TaskInput,
|
||||
Worker,
|
||||
WorkerStatus,
|
||||
)
|
||||
|
||||
from .base import UNSET, LightningStore, Unset
|
||||
from .base import UNSET, LightningStore, LightningStoreCapabilities, Unset
|
||||
|
||||
|
||||
class LightningStoreThreaded(LightningStore):
|
||||
@@ -35,6 +37,16 @@ class LightningStoreThreaded(LightningStore):
|
||||
self.store = store
|
||||
self._lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def capabilities(self) -> LightningStoreCapabilities:
|
||||
"""Return the capabilities of the store."""
|
||||
capabilities = self.store.capabilities
|
||||
return {
|
||||
**capabilities,
|
||||
"async_safe": True,
|
||||
"thread_safe": True,
|
||||
}
|
||||
|
||||
async def start_rollout(
|
||||
self,
|
||||
input: TaskInput,
|
||||
@@ -57,9 +69,9 @@ class LightningStoreThreaded(LightningStore):
|
||||
with self._lock:
|
||||
return await self.store.enqueue_rollout(input, mode, resources_id, config, metadata)
|
||||
|
||||
async def dequeue_rollout(self) -> Optional[AttemptedRollout]:
|
||||
async def dequeue_rollout(self, worker_id: Optional[str] = None) -> Optional[AttemptedRollout]:
|
||||
with self._lock:
|
||||
return await self.store.dequeue_rollout()
|
||||
return await self.store.dequeue_rollout(worker_id=worker_id)
|
||||
|
||||
async def start_attempt(self, rollout_id: str) -> AttemptedRollout:
|
||||
with self._lock:
|
||||
@@ -68,15 +80,48 @@ class LightningStoreThreaded(LightningStore):
|
||||
async def query_rollouts(
|
||||
self,
|
||||
*,
|
||||
status_in: Optional[Sequence[RolloutStatus]] = None,
|
||||
rollout_id_in: Optional[Sequence[str]] = None,
|
||||
rollout_id_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
status: Optional[Sequence[RolloutStatus]] = None,
|
||||
rollout_ids: Optional[Sequence[str]] = None,
|
||||
) -> List[Rollout]:
|
||||
) -> Sequence[Rollout]:
|
||||
with self._lock:
|
||||
return await self.store.query_rollouts(status=status, rollout_ids=rollout_ids)
|
||||
return await self.store.query_rollouts(
|
||||
status_in=status_in,
|
||||
rollout_id_in=rollout_id_in,
|
||||
rollout_id_contains=rollout_id_contains,
|
||||
filter_logic=filter_logic,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
status=status,
|
||||
rollout_ids=rollout_ids,
|
||||
)
|
||||
|
||||
async def query_attempts(self, rollout_id: str) -> List[Attempt]:
|
||||
async def query_attempts(
|
||||
self,
|
||||
rollout_id: str,
|
||||
*,
|
||||
sort_by: Optional[str] = "sequence_id",
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[Attempt]:
|
||||
with self._lock:
|
||||
return await self.store.query_attempts(rollout_id)
|
||||
return await self.store.query_attempts(
|
||||
rollout_id,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
async def get_rollout_by_id(self, rollout_id: str) -> Optional[Rollout]:
|
||||
with self._lock:
|
||||
@@ -86,6 +131,26 @@ class LightningStoreThreaded(LightningStore):
|
||||
with self._lock:
|
||||
return await self.store.get_latest_attempt(rollout_id)
|
||||
|
||||
async def query_resources(
|
||||
self,
|
||||
*,
|
||||
resources_id: Optional[str] = None,
|
||||
resources_id_contains: Optional[str] = None,
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[ResourcesUpdate]:
|
||||
with self._lock:
|
||||
return await self.store.query_resources(
|
||||
resources_id=resources_id,
|
||||
resources_id_contains=resources_id_contains,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
async def add_resources(self, resources: NamedResources) -> ResourcesUpdate:
|
||||
with self._lock:
|
||||
return await self.store.add_resources(resources)
|
||||
@@ -128,9 +193,39 @@ class LightningStoreThreaded(LightningStore):
|
||||
self,
|
||||
rollout_id: str,
|
||||
attempt_id: str | Literal["latest"] | None = None,
|
||||
) -> List[Span]:
|
||||
*,
|
||||
trace_id: Optional[str] = None,
|
||||
trace_id_contains: Optional[str] = None,
|
||||
span_id: Optional[str] = None,
|
||||
span_id_contains: Optional[str] = None,
|
||||
parent_id: Optional[str] = None,
|
||||
parent_id_contains: Optional[str] = None,
|
||||
name: Optional[str] = None,
|
||||
name_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
sort_by: Optional[str] = "sequence_id",
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
) -> Sequence[Span]:
|
||||
with self._lock:
|
||||
return await self.store.query_spans(rollout_id, attempt_id)
|
||||
return await self.store.query_spans(
|
||||
rollout_id,
|
||||
attempt_id,
|
||||
trace_id=trace_id,
|
||||
trace_id_contains=trace_id_contains,
|
||||
span_id=span_id,
|
||||
span_id_contains=span_id_contains,
|
||||
parent_id=parent_id,
|
||||
parent_id_contains=parent_id_contains,
|
||||
name=name,
|
||||
name_contains=name_contains,
|
||||
filter_logic=filter_logic,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
)
|
||||
|
||||
async def update_rollout(
|
||||
self,
|
||||
@@ -171,3 +266,39 @@ class LightningStoreThreaded(LightningStore):
|
||||
last_heartbeat_time=last_heartbeat_time,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
async def query_workers(
|
||||
self,
|
||||
*,
|
||||
status_in: Optional[Sequence[WorkerStatus]] = None,
|
||||
worker_id_contains: Optional[str] = None,
|
||||
filter_logic: Literal["and", "or"] = "and",
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: Literal["asc", "desc"] = "asc",
|
||||
limit: int = -1,
|
||||
offset: int = 0,
|
||||
) -> Sequence[Worker]:
|
||||
with self._lock:
|
||||
return await self.store.query_workers(
|
||||
status_in=status_in,
|
||||
worker_id_contains=worker_id_contains,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
async def get_worker_by_id(self, worker_id: str) -> Optional[Worker]:
|
||||
with self._lock:
|
||||
return await self.store.get_worker_by_id(worker_id)
|
||||
|
||||
async def update_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
heartbeat_stats: Dict[str, Any] | Unset = UNSET,
|
||||
) -> Worker:
|
||||
with self._lock:
|
||||
return await self.store.update_worker(
|
||||
worker_id=worker_id,
|
||||
heartbeat_stats=heartbeat_stats,
|
||||
)
|
||||
|
||||
@@ -2,25 +2,23 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from typing import TYPE_CHECKING, Any, AsyncGenerator, Awaitable, Iterator, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, AsyncGenerator, Iterator, List, Optional
|
||||
|
||||
import agentops
|
||||
import agentops.sdk.core
|
||||
import opentelemetry.trace as trace_api
|
||||
from agentops.sdk.core import TracingCore
|
||||
from agentops.sdk.processors import SpanProcessor
|
||||
from opentelemetry.instrumentation.utils import suppress_instrumentation
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from opentelemetry.sdk.trace import TracerProvider as TracerProviderImpl
|
||||
from opentelemetry.trace import get_tracer_provider
|
||||
from opentelemetry.trace.status import StatusCode
|
||||
|
||||
from agentlightning.instrumentation import instrument_all, uninstrument_all
|
||||
from agentlightning.instrumentation.agentops import AgentOpsServerManager
|
||||
from agentlightning.store.base import LightningStore
|
||||
|
||||
from .base import Tracer
|
||||
from .otel import LightningSpanProcessor, OtelTracer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentops.integration.callbacks.langchain import LangchainCallbackHandler
|
||||
@@ -29,7 +27,7 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AgentOpsTracer(Tracer):
|
||||
class AgentOpsTracer(OtelTracer):
|
||||
"""Traces agent execution using AgentOps.
|
||||
|
||||
This tracer provides functionality to capture execution details using the
|
||||
@@ -56,95 +54,37 @@ class AgentOpsTracer(Tracer):
|
||||
self.instrument_managed = instrument_managed
|
||||
self.daemon = daemon
|
||||
|
||||
self._agentops_server_manager = AgentOpsServerManager(self.daemon)
|
||||
self._agentops_server_port_val: Optional[int] = None
|
||||
|
||||
if not self.agentops_managed:
|
||||
logger.warning("agentops_managed=False. You are responsible for AgentOps setup.")
|
||||
if not self.instrument_managed:
|
||||
logger.warning("instrument_managed=False. You are responsible for all instrumentation.")
|
||||
|
||||
def __getstate__(self):
|
||||
state = self.__dict__.copy()
|
||||
state["_agentops_server_manager"] = None # Exclude the unpicklable server manager
|
||||
# _agentops_server_port_val (int) is inherently picklable and will be included.
|
||||
logger.debug(f"Getting state for pickling Trainer (PID {os.getpid()}). _agentops_server_manager excluded.")
|
||||
return state
|
||||
|
||||
def __setstate__(self, state: Any):
|
||||
self.__dict__.update(state)
|
||||
# In child process, self._agentops_server_manager will be None.
|
||||
logger.debug(f"Setting state for unpickled Trainer (PID {os.getpid()}). _agentops_server_manager is None.")
|
||||
|
||||
def init(self, *args: Any, **kwargs: Any):
|
||||
if self.agentops_managed and self._agentops_server_manager:
|
||||
self._agentops_server_manager.start()
|
||||
self._agentops_server_port_val = self._agentops_server_manager.get_port()
|
||||
if self._agentops_server_port_val is None:
|
||||
if (
|
||||
self._agentops_server_manager.server_process is not None
|
||||
and self._agentops_server_manager.server_process.is_alive()
|
||||
):
|
||||
raise RuntimeError("AgentOps server started but port is None. Check server manager logic.")
|
||||
elif (
|
||||
self._agentops_server_port_val is None and self._agentops_server_manager.server_process is None
|
||||
): # Server failed to start
|
||||
raise RuntimeError("AgentOps server manager indicates server is not running and port is None.")
|
||||
|
||||
def teardown(self):
|
||||
if self.agentops_managed:
|
||||
self._agentops_server_manager.stop()
|
||||
logger.info("AgentOps server stopped.")
|
||||
|
||||
def instrument(self, worker_id: int):
|
||||
instrument_all()
|
||||
|
||||
def uninstrument(self, worker_id: int):
|
||||
uninstrument_all()
|
||||
|
||||
def init_worker(self, worker_id: int):
|
||||
super().init_worker(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Setting up tracer...") # worker_id included in process name
|
||||
def _initialize_tracer_provider(self, worker_id: int):
|
||||
logger.info(f"[Worker {worker_id}] Setting up AgentOps tracer...") # worker_id included in process name
|
||||
|
||||
if self.instrument_managed:
|
||||
self.instrument(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Instrumentation applied.")
|
||||
|
||||
if self.agentops_managed:
|
||||
if self._agentops_server_port_val: # Use the stored, picklable port value
|
||||
base_url = f"http://localhost:{self._agentops_server_port_val}"
|
||||
env_vars_to_set = {
|
||||
"AGENTOPS_API_KEY": "dummy",
|
||||
"AGENTOPS_API_ENDPOINT": base_url,
|
||||
"AGENTOPS_APP_URL": f"{base_url}/notavailable",
|
||||
"AGENTOPS_EXPORTER_ENDPOINT": f"{base_url}/traces",
|
||||
}
|
||||
for key, value in env_vars_to_set.items():
|
||||
os.environ[key] = value
|
||||
logger.info(f"[Worker {worker_id}] Env var set: {key}={value}")
|
||||
else:
|
||||
logger.warning(
|
||||
f"[Worker {worker_id}] AgentOps managed, but local server port is not available. Client may not connect as expected."
|
||||
)
|
||||
|
||||
os.environ.setdefault("AGENTOPS_API_KEY", "dummy")
|
||||
if not agentops.get_client().initialized:
|
||||
agentops.init() # type: ignore
|
||||
agentops.init(auto_start_session=False) # type: ignore
|
||||
logger.info(f"[Worker {worker_id}] AgentOps client initialized.")
|
||||
else:
|
||||
logger.warning(f"[Worker {worker_id}] AgentOps client was already initialized.")
|
||||
|
||||
self._lightning_span_processor = LightningSpanProcessor()
|
||||
|
||||
try:
|
||||
# new versions
|
||||
instance = agentops.sdk.core.tracer
|
||||
# TODO: The span processor cannot be deleted once added.
|
||||
# This might be a problem if the tracer is entered and exited multiple times.
|
||||
instance.provider.add_span_processor(self._lightning_span_processor) # type: ignore
|
||||
except AttributeError:
|
||||
# old versions
|
||||
instance = TracingCore.get_instance() # type: ignore
|
||||
instance._provider.add_span_processor(self._lightning_span_processor) # type: ignore
|
||||
# TODO: The span processor cannot be deleted once added.
|
||||
# This might be a problem if the tracer is entered and exited multiple times.
|
||||
self._get_tracer_provider().add_span_processor(self._lightning_span_processor) # type: ignore
|
||||
|
||||
def teardown_worker(self, worker_id: int) -> None:
|
||||
super().teardown_worker(worker_id)
|
||||
@@ -161,7 +101,7 @@ class AgentOpsTracer(Tracer):
|
||||
store: Optional[LightningStore] = None,
|
||||
rollout_id: Optional[str] = None,
|
||||
attempt_id: Optional[str] = None,
|
||||
) -> AsyncGenerator[LightningSpanProcessor, None]:
|
||||
) -> AsyncGenerator[trace_api.Tracer, None]:
|
||||
"""
|
||||
Starts a new tracing context. This should be used as a context manager.
|
||||
|
||||
@@ -172,12 +112,10 @@ class AgentOpsTracer(Tracer):
|
||||
attempt_id: Optional attempt ID to add the spans to.
|
||||
|
||||
Yields:
|
||||
The [`LightningSpanProcessor`][agentlightning.tracer.agentops.LightningSpanProcessor] instance to collect spans.
|
||||
The OpenTelemetry tracer instance to collect spans.
|
||||
"""
|
||||
with self._trace_context_sync(
|
||||
name=name, store=store, rollout_id=rollout_id, attempt_id=attempt_id
|
||||
) as processor:
|
||||
yield processor
|
||||
with self._trace_context_sync(name=name, store=store, rollout_id=rollout_id, attempt_id=attempt_id) as tracer:
|
||||
yield tracer
|
||||
|
||||
@contextmanager
|
||||
def _trace_context_sync(
|
||||
@@ -187,31 +125,50 @@ class AgentOpsTracer(Tracer):
|
||||
store: Optional[LightningStore] = None,
|
||||
rollout_id: Optional[str] = None,
|
||||
attempt_id: Optional[str] = None,
|
||||
) -> Iterator[LightningSpanProcessor]:
|
||||
) -> Iterator[trace_api.Tracer]:
|
||||
"""Implementation of `trace_context` for synchronous execution."""
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
tracer_provider = self._get_tracer_provider()
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if name is not None:
|
||||
kwargs["trace_name"] = name
|
||||
elif rollout_id is not None:
|
||||
kwargs["trace_name"] = rollout_id
|
||||
if store is not None and rollout_id is not None and attempt_id is not None:
|
||||
if store.capabilities.get("otlp_traces", False) is True:
|
||||
logger.debug(f"Tracing to LightningStore rollout_id={rollout_id}, attempt_id={attempt_id}")
|
||||
self._enable_native_otlp_exporter(store, rollout_id, attempt_id)
|
||||
else:
|
||||
self._disable_native_otlp_exporter()
|
||||
ctx = self._lightning_span_processor.with_context(store=store, rollout_id=rollout_id, attempt_id=attempt_id)
|
||||
with ctx as processor:
|
||||
yield processor
|
||||
with ctx:
|
||||
# AgentOps end_trace and start_trace must live inside the lightning span processor context.
|
||||
# Otherwise some traces might not be recorded.
|
||||
with self._agentops_trace_context(rollout_id, attempt_id, kwargs):
|
||||
yield trace_api.get_tracer(__name__, tracer_provider=tracer_provider)
|
||||
elif store is None and rollout_id is None and attempt_id is None:
|
||||
# TODO: Add tests to cover both paths
|
||||
self._disable_native_otlp_exporter()
|
||||
with self._lightning_span_processor:
|
||||
yield self._lightning_span_processor
|
||||
with self._agentops_trace_context(None, None, kwargs):
|
||||
yield trace_api.get_tracer(__name__, tracer_provider=tracer_provider)
|
||||
else:
|
||||
raise ValueError("store, rollout_id, and attempt_id must be either all provided or all None")
|
||||
|
||||
def get_last_trace(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Retrieves the raw list of captured spans from the most recent trace.
|
||||
|
||||
Returns:
|
||||
A list of OpenTelemetry `ReadableSpan` objects.
|
||||
"""
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
return self._lightning_span_processor.spans()
|
||||
@contextmanager
|
||||
def _agentops_trace_context(self, rollout_id: Optional[str], attempt_id: Optional[str], kwargs: dict[str, Any]):
|
||||
trace = agentops.start_trace(**kwargs)
|
||||
status = StatusCode.OK # type: ignore
|
||||
try:
|
||||
yield
|
||||
except Exception as e:
|
||||
# TODO: I'm not sure whether this will catch errors in user code.
|
||||
status = StatusCode.ERROR # type: ignore
|
||||
logger.error(f"Trace failed for rollout_id={rollout_id}, attempt_id={attempt_id}: {e}")
|
||||
finally:
|
||||
agentops.end_trace(trace, end_state=status) # type: ignore
|
||||
|
||||
def get_langchain_handler(self, tags: List[str] | None = None) -> LangchainCallbackHandler:
|
||||
"""
|
||||
@@ -239,135 +196,26 @@ class AgentOpsTracer(Tracer):
|
||||
|
||||
get_langchain_callback_handler = get_langchain_handler # alias
|
||||
|
||||
def _get_tracer_provider(self) -> TracerProviderImpl:
|
||||
try:
|
||||
# new versions
|
||||
instance = agentops.sdk.core.tracer
|
||||
if instance.provider is None:
|
||||
raise RuntimeError("AgentOps TracerProvider is not initialized.")
|
||||
|
||||
class LightningSpanProcessor(SpanProcessor):
|
||||
"""Span processor that subclasses OpenTelemetry's `SpanProcessor` and adds support to dump traces
|
||||
to a [`LightningStore`][agentlightning.LightningStore].
|
||||
"""
|
||||
if get_tracer_provider() is not instance.provider:
|
||||
logger.error(
|
||||
"Mismatch between global singleton TracerProvider and AgentOps TracerProvider. "
|
||||
"AgentOps might not work properly."
|
||||
)
|
||||
|
||||
def __init__(self):
|
||||
self._spans: List[ReadableSpan] = []
|
||||
if not isinstance(instance.provider, TracerProviderImpl): # type: ignore
|
||||
raise RuntimeError("Unsupported TracerProvider type for AgentOps instrumentation.")
|
||||
|
||||
# Store related context and states
|
||||
self._store: Optional[LightningStore] = None
|
||||
self._rollout_id: Optional[str] = None
|
||||
self._attempt_id: Optional[str] = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# private asyncio loop running in a daemon thread
|
||||
self._loop_ready = threading.Event()
|
||||
self._loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
self._loop_thread = threading.Thread(target=self._loop_runner, name="otel-loop", daemon=True)
|
||||
self._loop_thread.start()
|
||||
self._loop_ready.wait() # loop is ready
|
||||
|
||||
def _loop_runner(self):
|
||||
loop = asyncio.new_event_loop()
|
||||
self._loop = loop
|
||||
asyncio.set_event_loop(loop)
|
||||
self._loop_ready.set()
|
||||
loop.run_forever()
|
||||
loop.close()
|
||||
|
||||
def __enter__(self):
|
||||
self._last_trace = None
|
||||
self._spans = []
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any):
|
||||
self._store = None
|
||||
self._rollout_id = None
|
||||
self._attempt_id = None
|
||||
|
||||
def _await_in_loop(self, coro: Awaitable[Any], timeout: Optional[float] = None) -> Any:
|
||||
# submit to the dedicated loop and wait synchronously
|
||||
if self._loop is None:
|
||||
raise RuntimeError("Loop is not initialized. This should not happen.")
|
||||
|
||||
# If already on the exporter loop thread, schedule and return immediately.
|
||||
# ---------------------------------------------------------------------------
|
||||
# WHY THIS CONDITIONAL EXISTS:
|
||||
# In rare cases, span.end() is triggered from a LangchainCallbackHandler.__del__
|
||||
# (or another finalizer) while the Python garbage collector is running on the
|
||||
# *same thread* that owns our exporter event loop ("otel-loop").
|
||||
#
|
||||
# When that happens, on_end() executes on the exporter loop thread itself.
|
||||
# If we were to call `asyncio.run_coroutine_threadsafe(...).result()` here,
|
||||
# it would deadlock immediately — because the loop cannot both wait on and run
|
||||
# the same coroutine. The Future stays pending forever and the loop stops
|
||||
# processing scheduled callbacks.
|
||||
#
|
||||
# To avoid that self-deadlock, we detect when on_end() runs on the exporter
|
||||
# loop thread. If so, we *schedule* the coroutine on the loop (fire-and-forget)
|
||||
# instead of blocking with .result().
|
||||
#
|
||||
# This situation can occur because Python calls __del__ in whatever thread
|
||||
# releases the last reference, which can easily be our loop thread if the
|
||||
# object is dereferenced during loop._run_once().
|
||||
# ---------------------------------------------------------------------------
|
||||
if threading.current_thread() is self._loop_thread:
|
||||
self._loop.call_soon_threadsafe(asyncio.create_task, coro) # type: ignore
|
||||
return None
|
||||
|
||||
fut = asyncio.run_coroutine_threadsafe(coro, self._loop) # type: ignore
|
||||
return fut.result(timeout=timeout) # raises on error # type: ignore
|
||||
|
||||
def shutdown(self) -> None:
|
||||
if self._loop:
|
||||
self._loop.call_soon_threadsafe(self._loop.stop)
|
||||
self._loop_thread.join(timeout=5)
|
||||
self._loop = None
|
||||
|
||||
def force_flush(self, timeout_millis: int = 30000) -> bool:
|
||||
return True
|
||||
|
||||
def spans(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Get the list of spans collected by this processor.
|
||||
This is useful for debugging and testing purposes.
|
||||
|
||||
Returns:
|
||||
List of ReadableSpan objects collected during tracing.
|
||||
"""
|
||||
return self._spans
|
||||
|
||||
def with_context(self, store: LightningStore, rollout_id: str, attempt_id: str):
|
||||
# simple context manager without nesting into asyncio
|
||||
class _Ctx:
|
||||
def __enter__(_): # type: ignore
|
||||
with self._lock:
|
||||
self._store, self._rollout_id, self._attempt_id = store, rollout_id, attempt_id
|
||||
self._last_trace = None
|
||||
self._spans = []
|
||||
return self
|
||||
|
||||
def __exit__(_, exc_type, exc, tb): # type: ignore
|
||||
with self._lock:
|
||||
self._store = self._rollout_id = self._attempt_id = None
|
||||
|
||||
return _Ctx()
|
||||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
"""
|
||||
Process a span when it ends.
|
||||
|
||||
Args:
|
||||
span: The span that has ended.
|
||||
"""
|
||||
# Skip if span is not sampled
|
||||
if not span.context or not span.context.trace_flags.sampled:
|
||||
return
|
||||
|
||||
if self._store and self._rollout_id and self._attempt_id:
|
||||
try:
|
||||
# Submit add_otel_span to the event loop and wait for it to complete
|
||||
with suppress_instrumentation():
|
||||
self._await_in_loop(
|
||||
self._store.add_otel_span(self._rollout_id, self._attempt_id, span),
|
||||
timeout=60.0,
|
||||
)
|
||||
except Exception:
|
||||
# log; on_end MUST NOT raise
|
||||
logger.exception(f"Error adding span to store: {span.name}")
|
||||
|
||||
self._spans.append(span)
|
||||
self._tracer_provider = instance.provider
|
||||
return self._tracer_provider
|
||||
except AttributeError:
|
||||
# old versions
|
||||
instance = TracingCore.get_instance() # type: ignore
|
||||
self._tracer_provider = instance._provider # type: ignore
|
||||
return self._tracer_provider # type: ignore
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Any, AsyncContextManager, Awaitable, Callable, ContextManager, List, Optional
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
@@ -138,3 +139,27 @@ class Tracer(ParallelWorkerBase):
|
||||
"""
|
||||
logger.warning(f"{self.__class__.__name__} does not provide a LangChain callback handler.")
|
||||
return None
|
||||
|
||||
@contextmanager
|
||||
def lifespan(self):
|
||||
"""A context manager to manage the lifespan of the tracer.
|
||||
|
||||
This can be used to set up and tear down any necessary resources
|
||||
for the tracer, useful for debugging purposes.
|
||||
"""
|
||||
has_init = False
|
||||
has_init_worker = False
|
||||
try:
|
||||
self.init()
|
||||
has_init = True
|
||||
|
||||
self.init_worker(0)
|
||||
has_init_worker = True
|
||||
|
||||
yield
|
||||
|
||||
finally:
|
||||
if has_init_worker:
|
||||
self.teardown_worker(0)
|
||||
if has_init:
|
||||
self.teardown()
|
||||
|
||||
+255
-11
@@ -2,16 +2,25 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncGenerator, List, Optional
|
||||
from typing import Any, AsyncGenerator, Awaitable, List, Optional
|
||||
|
||||
import opentelemetry.trace as trace_api
|
||||
from opentelemetry.sdk.trace import ReadableSpan, TracerProvider
|
||||
from agentops.sdk.core import BatchSpanProcessor
|
||||
from opentelemetry.instrumentation.utils import suppress_instrumentation
|
||||
from opentelemetry.sdk.resources import Resource
|
||||
from opentelemetry.sdk.trace import ReadableSpan, SpanProcessor
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
from opentelemetry.sdk.trace import TracerProvider as TracerProviderImpl
|
||||
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
|
||||
|
||||
from agentlightning.store.base import LightningStore
|
||||
from agentlightning.types.tracer import SpanNames
|
||||
from agentlightning.utils.otlp import LightningStoreOTLPExporter
|
||||
|
||||
from .agentops import LightningSpanProcessor # FIXME: This import should be from otel to agentops
|
||||
from .base import Tracer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -29,21 +38,31 @@ class OtelTracer(Tracer):
|
||||
# This provider is only initialized when the worker is initialized.
|
||||
self._tracer_provider: Optional[TracerProvider] = None
|
||||
self._lightning_span_processor: Optional[LightningSpanProcessor] = None
|
||||
self._simple_span_processor: Optional[SimpleSpanProcessor] = None
|
||||
self._otlp_span_exporter: Optional[LightningStoreOTLPExporter] = None
|
||||
self._initialized: bool = False
|
||||
|
||||
def init_worker(self, worker_id: int):
|
||||
super().init_worker(worker_id)
|
||||
self._initialize_tracer_provider(worker_id)
|
||||
|
||||
def _initialize_tracer_provider(self, worker_id: int):
|
||||
logger.info(f"[Worker {worker_id}] Setting up OpenTelemetry tracer...")
|
||||
|
||||
if self._initialized:
|
||||
logger.error("Tracer provider is already initialized. OpenTelemetry may not work as expected.")
|
||||
|
||||
tracer_provider = TracerProvider()
|
||||
trace_api.set_tracer_provider(tracer_provider)
|
||||
self._tracer_provider = TracerProvider()
|
||||
trace_api.set_tracer_provider(self._tracer_provider)
|
||||
self._lightning_span_processor = LightningSpanProcessor()
|
||||
tracer_provider.add_span_processor(self._lightning_span_processor)
|
||||
self._tracer_provider.add_span_processor(self._lightning_span_processor)
|
||||
self._otlp_span_exporter = LightningStoreOTLPExporter()
|
||||
self._simple_span_processor = SimpleSpanProcessor(self._otlp_span_exporter)
|
||||
self._tracer_provider.add_span_processor(self._simple_span_processor)
|
||||
self._initialized = True
|
||||
|
||||
logger.info(f"[Worker {worker_id}] OpenTelemetry tracer provider initialized.")
|
||||
|
||||
def teardown_worker(self, worker_id: int):
|
||||
super().teardown_worker(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Tearing down OpenTelemetry tracer...")
|
||||
@@ -57,7 +76,7 @@ class OtelTracer(Tracer):
|
||||
store: Optional[LightningStore] = None,
|
||||
rollout_id: Optional[str] = None,
|
||||
attempt_id: Optional[str] = None,
|
||||
) -> AsyncGenerator[LightningSpanProcessor, None]:
|
||||
) -> AsyncGenerator[trace_api.Tracer, None]:
|
||||
"""
|
||||
Starts a new tracing context. This should be used as a context manager.
|
||||
|
||||
@@ -68,18 +87,24 @@ class OtelTracer(Tracer):
|
||||
attempt_id: Optional attempt ID to add the spans to.
|
||||
|
||||
Yields:
|
||||
The LightningSpanProcessor instance to collect spans.
|
||||
The OpenTelemetry tracer instance to collect spans.
|
||||
"""
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
|
||||
if store is not None and rollout_id is not None and attempt_id is not None:
|
||||
if store.capabilities.get("otlp_traces", False) is True:
|
||||
logger.debug(f"Tracing to LightningStore rollout_id={rollout_id}, attempt_id={attempt_id}")
|
||||
self._enable_native_otlp_exporter(store, rollout_id, attempt_id)
|
||||
else:
|
||||
self._disable_native_otlp_exporter()
|
||||
ctx = self._lightning_span_processor.with_context(store=store, rollout_id=rollout_id, attempt_id=attempt_id)
|
||||
with ctx as processor:
|
||||
yield processor
|
||||
with ctx:
|
||||
yield trace_api.get_tracer(__name__, tracer_provider=self._tracer_provider)
|
||||
elif store is None and rollout_id is None and attempt_id is None:
|
||||
self._disable_native_otlp_exporter()
|
||||
with self._lightning_span_processor:
|
||||
yield self._lightning_span_processor
|
||||
yield trace_api.get_tracer(__name__, tracer_provider=self._tracer_provider)
|
||||
else:
|
||||
raise ValueError("store, rollout_id, and attempt_id must be either all provided or all None")
|
||||
|
||||
@@ -93,3 +118,222 @@ class OtelTracer(Tracer):
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
return self._lightning_span_processor.spans()
|
||||
|
||||
def _get_tracer_provider(self) -> TracerProviderImpl:
|
||||
if self._tracer_provider is None:
|
||||
raise RuntimeError("TracerProvider is not initialized. Call init_worker() first.")
|
||||
return self._tracer_provider
|
||||
|
||||
def _enable_native_otlp_exporter(self, store: LightningStore, rollout_id: str, attempt_id: str):
|
||||
tracer_provider = self._get_tracer_provider()
|
||||
active_span_processor = tracer_provider._active_span_processor # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
# Override the resources so that the server knows where the request comes from.
|
||||
tracer_provider._resource = tracer_provider._resource.merge( # pyright: ignore[reportPrivateUsage]
|
||||
Resource.create(
|
||||
{
|
||||
SpanNames.ROLLOUT_ID: rollout_id,
|
||||
SpanNames.ATTEMPT_ID: attempt_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
instrumented = False
|
||||
candidates: List[str] = []
|
||||
for processor in active_span_processor._span_processors: # pyright: ignore[reportPrivateUsage]
|
||||
if isinstance(processor, LightningSpanProcessor):
|
||||
# We don't need the LightningSpanProcessor any more.
|
||||
logger.debug("LightningSpanProcessor already present in TracerProvider, disabling it.")
|
||||
processor.disable_store_submission = True
|
||||
elif isinstance(processor, (SimpleSpanProcessor, BatchSpanProcessor)):
|
||||
# Instead, we rely on the OTLPSpanExporter to send spans to the store.
|
||||
if isinstance(processor.span_exporter, LightningStoreOTLPExporter):
|
||||
processor.span_exporter.enable_store_otlp(store.otlp_traces_endpoint(), rollout_id, attempt_id)
|
||||
logger.debug(f"Set LightningStoreOTLPExporter endpoint to {store.otlp_traces_endpoint()}")
|
||||
instrumented = True
|
||||
else:
|
||||
candidates.append(
|
||||
f"{processor.__class__.__name__} with {processor.span_exporter.__class__.__name__}"
|
||||
)
|
||||
else:
|
||||
candidates.append(f"{processor.__class__.__name__}")
|
||||
|
||||
if not instrumented:
|
||||
raise RuntimeError(
|
||||
"Failed to enable native OTLP exporter: no BatchSpanProcessor or SimpleSpanProcessor with "
|
||||
"LightningStoreOTLPExporter found in TracerProvider. Please try using a non-OTLP store."
|
||||
"Candidates are: " + ", ".join(candidates)
|
||||
)
|
||||
|
||||
def _disable_native_otlp_exporter(self):
|
||||
tracer_provider = self._get_tracer_provider()
|
||||
active_span_processor = tracer_provider._active_span_processor # pyright: ignore[reportPrivateUsage]
|
||||
tracer_provider._resource = tracer_provider._resource.merge( # pyright: ignore[reportPrivateUsage]
|
||||
Resource.create(
|
||||
{
|
||||
SpanNames.ROLLOUT_ID: "",
|
||||
SpanNames.ATTEMPT_ID: "",
|
||||
}
|
||||
)
|
||||
) # reset resource
|
||||
for processor in active_span_processor._span_processors: # pyright: ignore[reportPrivateUsage]
|
||||
if isinstance(processor, LightningSpanProcessor):
|
||||
# We will be in need of the LightningSpanProcessor again.
|
||||
logger.debug("Enabling LightningSpanProcessor in TracerProvider.")
|
||||
processor.disable_store_submission = False
|
||||
|
||||
|
||||
class LightningSpanProcessor(SpanProcessor):
|
||||
"""Span processor that subclasses OpenTelemetry's `SpanProcessor` and adds support to dump traces
|
||||
to a [`LightningStore`][agentlightning.LightningStore].
|
||||
|
||||
It serves two purposes:
|
||||
|
||||
1. Records all the spans in a local buffer.
|
||||
2. Submits the spans to the event loop to be added to the store.
|
||||
"""
|
||||
|
||||
def __init__(self, disable_store_submission: bool = False):
|
||||
self._disable_store_submission: bool = disable_store_submission
|
||||
self._spans: List[ReadableSpan] = []
|
||||
|
||||
# Store related context and states
|
||||
self._store: Optional[LightningStore] = None
|
||||
self._rollout_id: Optional[str] = None
|
||||
self._attempt_id: Optional[str] = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# private asyncio loop running in a daemon thread
|
||||
self._loop_ready = threading.Event()
|
||||
self._loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
|
||||
@property
|
||||
def disable_store_submission(self) -> bool:
|
||||
"""Whether to disable submitting spans to the store."""
|
||||
return self._disable_store_submission
|
||||
|
||||
@disable_store_submission.setter
|
||||
def disable_store_submission(self, value: bool) -> None:
|
||||
self._disable_store_submission = value
|
||||
|
||||
def _ensure_loop(self) -> None:
|
||||
if self._loop_thread is None or self._loop is None:
|
||||
self._loop_ready.clear()
|
||||
self._loop_thread = threading.Thread(target=self._loop_runner, name="otel-loop", daemon=True)
|
||||
self._loop_thread.start()
|
||||
self._loop_ready.wait() # loop is ready
|
||||
|
||||
def _loop_runner(self):
|
||||
loop = asyncio.new_event_loop()
|
||||
self._loop = loop
|
||||
asyncio.set_event_loop(loop)
|
||||
self._loop_ready.set()
|
||||
loop.run_forever()
|
||||
loop.close()
|
||||
|
||||
def __enter__(self):
|
||||
self._last_trace = None
|
||||
self._spans = []
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any):
|
||||
self._store = None
|
||||
self._rollout_id = None
|
||||
self._attempt_id = None
|
||||
|
||||
def _await_in_loop(self, coro: Awaitable[Any], timeout: Optional[float] = None) -> Any:
|
||||
# submit to the dedicated loop and wait synchronously
|
||||
self._ensure_loop()
|
||||
if self._loop is None:
|
||||
raise RuntimeError("Loop is not initialized. This should not happen.")
|
||||
|
||||
# If already on the exporter loop thread, schedule and return immediately.
|
||||
# ---------------------------------------------------------------------------
|
||||
# WHY THIS CONDITIONAL EXISTS:
|
||||
# In rare cases, span.end() is triggered from a LangchainCallbackHandler.__del__
|
||||
# (or another finalizer) while the Python garbage collector is running on the
|
||||
# *same thread* that owns our exporter event loop ("otel-loop").
|
||||
#
|
||||
# When that happens, on_end() executes on the exporter loop thread itself.
|
||||
# If we were to call `asyncio.run_coroutine_threadsafe(...).result()` here,
|
||||
# it would deadlock immediately — because the loop cannot both wait on and run
|
||||
# the same coroutine. The Future stays pending forever and the loop stops
|
||||
# processing scheduled callbacks.
|
||||
#
|
||||
# To avoid that self-deadlock, we detect when on_end() runs on the exporter
|
||||
# loop thread. If so, we *schedule* the coroutine on the loop (fire-and-forget)
|
||||
# instead of blocking with .result().
|
||||
#
|
||||
# This situation can occur because Python calls __del__ in whatever thread
|
||||
# releases the last reference, which can easily be our loop thread if the
|
||||
# object is dereferenced during loop._run_once().
|
||||
# ---------------------------------------------------------------------------
|
||||
if threading.current_thread() is self._loop_thread:
|
||||
self._loop.call_soon_threadsafe(asyncio.create_task, coro) # type: ignore
|
||||
return None
|
||||
|
||||
fut = asyncio.run_coroutine_threadsafe(coro, self._loop) # type: ignore
|
||||
return fut.result(timeout=timeout) # raises on error # type: ignore
|
||||
|
||||
def shutdown(self) -> None:
|
||||
if self._loop:
|
||||
self._loop.call_soon_threadsafe(self._loop.stop)
|
||||
self._loop = None
|
||||
if self._loop_thread:
|
||||
self._loop_thread.join(timeout=5)
|
||||
|
||||
def force_flush(self, timeout_millis: int = 30000) -> bool:
|
||||
return True
|
||||
|
||||
def spans(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Get the list of spans collected by this processor.
|
||||
This is useful for debugging and testing purposes.
|
||||
|
||||
Returns:
|
||||
List of ReadableSpan objects collected during tracing.
|
||||
"""
|
||||
return self._spans
|
||||
|
||||
def with_context(self, store: LightningStore, rollout_id: str, attempt_id: str):
|
||||
# simple context manager without nesting into asyncio
|
||||
class _Ctx:
|
||||
def __enter__(_): # type: ignore
|
||||
# Use _ instead of self to avoid shadowing the instance method.
|
||||
with self._lock:
|
||||
self._store, self._rollout_id, self._attempt_id = store, rollout_id, attempt_id
|
||||
self._last_trace = None
|
||||
self._spans = []
|
||||
return self
|
||||
|
||||
def __exit__(_, exc_type, exc, tb): # type: ignore
|
||||
with self._lock:
|
||||
self._store = self._rollout_id = self._attempt_id = None
|
||||
|
||||
return _Ctx()
|
||||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
"""
|
||||
Process a span when it ends.
|
||||
|
||||
Args:
|
||||
span: The span that has ended.
|
||||
"""
|
||||
# Skip if span is not sampled
|
||||
if not span.context or not span.context.trace_flags.sampled:
|
||||
return
|
||||
|
||||
if not self._disable_store_submission and self._store and self._rollout_id and self._attempt_id:
|
||||
try:
|
||||
# Submit add_otel_span to the event loop and wait for it to complete
|
||||
with suppress_instrumentation():
|
||||
self._ensure_loop()
|
||||
self._await_in_loop(
|
||||
self._store.add_otel_span(self._rollout_id, self._attempt_id, span),
|
||||
timeout=60.0,
|
||||
)
|
||||
except Exception:
|
||||
# log; on_end MUST NOT raise
|
||||
logger.exception(f"Error adding span to store: {span.name}")
|
||||
|
||||
self._spans.append(span)
|
||||
|
||||
@@ -10,14 +10,19 @@ from typing import (
|
||||
Callable,
|
||||
Dict,
|
||||
Generic,
|
||||
Iterator,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Protocol,
|
||||
Sequence,
|
||||
SupportsIndex,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
overload,
|
||||
)
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
@@ -49,6 +54,12 @@ __all__ = [
|
||||
"Attempt",
|
||||
"AttemptedRollout",
|
||||
"Hook",
|
||||
"Worker",
|
||||
"WorkerStatus",
|
||||
"PaginatedResult",
|
||||
"FilterOptions",
|
||||
"SortOptions",
|
||||
"FilterField",
|
||||
]
|
||||
|
||||
T_co = TypeVar("T_co", covariant=True)
|
||||
@@ -200,6 +211,32 @@ class AttemptedRollout(Rollout):
|
||||
return self
|
||||
|
||||
|
||||
WorkerStatus = Literal["idle", "busy", "unknown"]
|
||||
|
||||
|
||||
class Worker(BaseModel):
|
||||
"""Worker information. This is actually the same as Runner info."""
|
||||
|
||||
worker_id: str
|
||||
"""The ID of the worker."""
|
||||
status: WorkerStatus = "unknown"
|
||||
"""The status of the worker."""
|
||||
heartbeat_stats: Optional[Dict[str, Any]] = None
|
||||
"""Statistics about the worker's heartbeat."""
|
||||
last_heartbeat_time: Optional[float] = None
|
||||
"""The last time when the worker has reported the stats."""
|
||||
last_dequeue_time: Optional[float] = None
|
||||
"""The last time when the worker has tried to dequeue a rollout."""
|
||||
last_busy_time: Optional[float] = None
|
||||
"""The last time when the worker has started an attempt and became busy."""
|
||||
last_idle_time: Optional[float] = None
|
||||
"""The last time when the worker has triggered the end of an attempt and became idle."""
|
||||
current_rollout_id: Optional[str] = None
|
||||
"""The ID of the current rollout that the worker is processing."""
|
||||
current_attempt_id: Optional[str] = None
|
||||
"""The ID of the current attempt that the worker is processing."""
|
||||
|
||||
|
||||
TaskInput = Any
|
||||
"""Task input type. Accepts arbitrary payloads."""
|
||||
|
||||
@@ -393,3 +430,104 @@ class Hook(ParallelWorkerBase):
|
||||
Subclasses can override this method for cleanup or additional
|
||||
logging. By default, this is a no-op.
|
||||
"""
|
||||
|
||||
|
||||
class FilterField(TypedDict, total=False):
|
||||
"""An operator dict for a single field."""
|
||||
|
||||
exact: Any
|
||||
within: Sequence[Any]
|
||||
contains: str
|
||||
|
||||
|
||||
FilterOptions = Mapping[
|
||||
Union[str, Literal["_aggregate", "_must"]],
|
||||
Union[FilterField, Literal["and", "or"], Mapping[str, FilterField]],
|
||||
]
|
||||
"""A mapping of field name -> operator dict.
|
||||
|
||||
Each operator dict can contain:
|
||||
|
||||
- "exact": value for exact equality.
|
||||
- "within": iterable of allowed values.
|
||||
- "contains": substring to search for in string fields.
|
||||
|
||||
The filter can also have a special field called "_aggregate" that can be used to specify the logic
|
||||
to combine the results of the filters:
|
||||
|
||||
- "and": all conditions must match. This is the default value if not specified.
|
||||
- "or": at least one condition must match.
|
||||
|
||||
All conditions within a field and between different fields are
|
||||
stored in a unified pool and combined using `_aggregate`.
|
||||
|
||||
The filter can also have a special group called "_must", which is a mapping of filters that must all match,
|
||||
no matter whether the aggregate logic is "and" or "or".
|
||||
|
||||
Example:
|
||||
|
||||
```json
|
||||
{
|
||||
"_aggregate": "or",
|
||||
"_must": {
|
||||
"city": {"exact": "New York"},
|
||||
"timezone": {"within": ["America/New_York", "America/Los_Angeles"]},
|
||||
},
|
||||
"status": {"exact": "active"},
|
||||
"id": {"within": [1, 2, 3]},
|
||||
"name": {"contains": "foo"},
|
||||
}
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
class SortOptions(TypedDict):
|
||||
"""Options for sorting the collection."""
|
||||
|
||||
name: str
|
||||
"""The name of the field to sort by."""
|
||||
order: Literal["asc", "desc"]
|
||||
"""The order to sort by."""
|
||||
|
||||
|
||||
T_item = TypeVar("T_item")
|
||||
|
||||
|
||||
class PaginatedResult(BaseModel, Sequence[T_item]):
|
||||
"""Result of a paginated query.
|
||||
|
||||
Behaves like a sequence, but also carries pagination metadata (limit, offset, total).
|
||||
"""
|
||||
|
||||
items: Sequence[T_item]
|
||||
"""Items in the result."""
|
||||
limit: int
|
||||
"""Limit of the result."""
|
||||
offset: int
|
||||
"""Offset of the result."""
|
||||
total: int
|
||||
"""Total number of items in the collection."""
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.items)
|
||||
|
||||
@overload
|
||||
def __getitem__(self, index: int) -> T_item: ...
|
||||
|
||||
@overload
|
||||
def __getitem__(self, index: slice) -> Sequence[T_item]: ...
|
||||
|
||||
def __getitem__(self, index: Union[int, slice]) -> Union[T_item, Sequence[T_item]]:
|
||||
return self.items[index]
|
||||
|
||||
# Overriding __iter__ enables list(paginated_result) to work as expected,
|
||||
# but changes Pydantic's default dict iteration behavior (which would otherwise
|
||||
# iterate over field names).
|
||||
def __iter__(self) -> Iterator[T_item]: # type: ignore
|
||||
return iter(self.items)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
first_item_repr = repr(self.items[0]) if self.items else "empty"
|
||||
items_repr = f"[{first_item_repr}, ...]" if len(self.items) > 1 else first_item_repr
|
||||
slice_repr = f"{self.offset}:" if self.limit == -1 else f"{self.offset}:{self.offset + self.limit}"
|
||||
return f"<PaginatedResult ({slice_repr} of {self.total}) {items_repr}>"
|
||||
|
||||
@@ -194,5 +194,11 @@ class ResourcesUpdate(BaseModel):
|
||||
|
||||
resources_id: str
|
||||
"""Identifier used to version the resources."""
|
||||
create_time: float
|
||||
"""Timestamp of the creation time of the resources."""
|
||||
update_time: float
|
||||
"""Timestamp of the last update time of the resources."""
|
||||
version: int
|
||||
"""Version of the resources."""
|
||||
resources: NamedResources
|
||||
"""Mapping of resource names to their definitions."""
|
||||
|
||||
@@ -411,6 +411,12 @@ class SpanNames(str, Enum):
|
||||
"""The name of the exception span."""
|
||||
VIRTUAL = "agentlightning.virtual"
|
||||
"""The name of the virtual span. It represents derived spans without concrete operations."""
|
||||
ROLLOUT_ID = "agentlightning.rollout_id"
|
||||
"""The name of the rollout ID."""
|
||||
ATTEMPT_ID = "agentlightning.attempt_id"
|
||||
"""The name of the attempt ID."""
|
||||
SPAN_SEQUENCE_ID = "agentlightning.span_sequence_id"
|
||||
"""The name of the span sequence ID."""
|
||||
|
||||
|
||||
class SpanAttributeNames(str, Enum):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
@@ -0,0 +1,428 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import gzip
|
||||
import logging
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Optional, Sequence, Tuple, Type, TypeVar
|
||||
|
||||
from fastapi import Request, Response
|
||||
from google.protobuf import json_format
|
||||
from google.rpc.status_pb2 import Status
|
||||
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
||||
from opentelemetry.proto.collector.logs.v1.logs_service_pb2 import (
|
||||
ExportLogsServiceRequest,
|
||||
ExportLogsServiceResponse,
|
||||
)
|
||||
from opentelemetry.proto.collector.metrics.v1.metrics_service_pb2 import (
|
||||
ExportMetricsServiceRequest,
|
||||
ExportMetricsServiceResponse,
|
||||
)
|
||||
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import (
|
||||
ExportTraceServiceRequest,
|
||||
ExportTraceServiceResponse,
|
||||
)
|
||||
from opentelemetry.proto.common.v1.common_pb2 import AnyValue, KeyValue
|
||||
from opentelemetry.proto.resource.v1.resource_pb2 import Resource as ProtoResource
|
||||
from opentelemetry.proto.trace.v1.trace_pb2 import Span as ProtoSpan
|
||||
from opentelemetry.proto.trace.v1.trace_pb2 import Status as ProtoStatus
|
||||
from opentelemetry.sdk.resources import Resource
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from opentelemetry.sdk.trace.export import SpanExportResult
|
||||
from opentelemetry.util.types import AttributeValue
|
||||
|
||||
from agentlightning.store.base import LightningStore
|
||||
from agentlightning.types.tracer import (
|
||||
Attributes,
|
||||
Event,
|
||||
Link,
|
||||
OtelResource,
|
||||
Span,
|
||||
SpanContext,
|
||||
SpanNames,
|
||||
TraceStatus,
|
||||
convert_timestamp,
|
||||
)
|
||||
|
||||
PROTOBUF_CT = "application/x-protobuf"
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
T_request = TypeVar("T_request", ExportLogsServiceRequest, ExportMetricsServiceRequest, ExportTraceServiceRequest)
|
||||
T_response = TypeVar("T_response", ExportLogsServiceResponse, ExportMetricsServiceResponse, ExportTraceServiceResponse)
|
||||
|
||||
|
||||
async def handle_otlp_export(
|
||||
request: Request,
|
||||
request_message_cls: Type[T_request],
|
||||
response_message_cls: Type[T_response],
|
||||
message_callback: Optional[Callable[[T_request], Awaitable[None]]],
|
||||
signal_name: str,
|
||||
) -> Response:
|
||||
"""
|
||||
Generic handler for /v1/traces, /v1/metrics, /v1/logs.
|
||||
|
||||
Convert the OTLP Protobuf request to a JSON-like object.
|
||||
"""
|
||||
content_type = request.headers.get("Content-Type", "").split(";")[0].strip()
|
||||
|
||||
if content_type != PROTOBUF_CT:
|
||||
# For brevity we only support binary protobuf here.
|
||||
return _bad_request_response(
|
||||
request,
|
||||
f"Unsupported Content-Type '{content_type}', expected '{PROTOBUF_CT}'",
|
||||
content_type=PROTOBUF_CT,
|
||||
)
|
||||
|
||||
raw_body = await request.body()
|
||||
body = _read_body_maybe_gzip(request, raw_body)
|
||||
|
||||
# Empty request is allowed and should still succeed.
|
||||
if not body:
|
||||
req_msg = request_message_cls()
|
||||
else:
|
||||
req_msg = request_message_cls()
|
||||
try:
|
||||
req_msg.ParseFromString(body)
|
||||
except Exception as exc:
|
||||
return _bad_request_response(request, f"Unable to parse OTLP {signal_name} payload: {exc}")
|
||||
|
||||
if message_callback is not None:
|
||||
await message_callback(req_msg)
|
||||
|
||||
# Build success response. Partial success field is left unset.
|
||||
resp_msg = response_message_cls()
|
||||
|
||||
# Encode response in the same Content-Type as request.
|
||||
if content_type == PROTOBUF_CT:
|
||||
resp_bytes = resp_msg.SerializeToString()
|
||||
else:
|
||||
resp_bytes = json_format.MessageToJson(resp_msg).encode("utf-8")
|
||||
|
||||
resp_bytes, headers = _maybe_gzip_response(request, resp_bytes)
|
||||
|
||||
return Response(
|
||||
content=resp_bytes,
|
||||
media_type=content_type,
|
||||
status_code=200,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
async def spans_from_proto(request: ExportTraceServiceRequest, store: LightningStore) -> List[Span]:
|
||||
"""Parse an OTLP proto payload into List[Span].
|
||||
|
||||
A store is needed here for generating a sequence ID for each span.
|
||||
"""
|
||||
output_spans: List[Span] = []
|
||||
|
||||
for resource_spans in request.resource_spans:
|
||||
# Resource-level attributes & IDs
|
||||
resource_attrs = _kv_list_to_dict(resource_spans.resource.attributes)
|
||||
# rollout_id, attempt_id from resource attributes when present.
|
||||
rollout_id_resource = resource_attrs.get(SpanNames.ROLLOUT_ID)
|
||||
attempt_id_resource = resource_attrs.get(SpanNames.ATTEMPT_ID)
|
||||
# If sequence id is provided, all the spans will share the same sequence ID.
|
||||
# unless otherwise overridden by span-level attributes.
|
||||
sequence_id_resource = resource_attrs.get(SpanNames.SPAN_SEQUENCE_ID)
|
||||
|
||||
otel_resource = _resource_from_proto(resource_spans.resource, getattr(resource_spans, "schema_url", ""))
|
||||
|
||||
# Each ScopeSpans contains multiple spans
|
||||
for scope_spans in resource_spans.scope_spans:
|
||||
for proto_span in scope_spans.spans:
|
||||
trace_id_hex = _bytes_to_trace_id_hex(proto_span.trace_id)
|
||||
span_id_hex = _bytes_to_span_id_hex(proto_span.span_id)
|
||||
parent_id_hex = _bytes_to_span_id_hex(proto_span.parent_span_id) if proto_span.parent_span_id else None
|
||||
|
||||
# Status
|
||||
status_code_str = _STATUS_CODE_MAP.get(proto_span.status.code, "UNSET")
|
||||
status = TraceStatus(
|
||||
status_code=status_code_str,
|
||||
description=proto_span.status.message or None,
|
||||
)
|
||||
|
||||
# Attributes
|
||||
span_attrs = _kv_list_to_dict(proto_span.attributes)
|
||||
|
||||
# Context
|
||||
context = SpanContext(
|
||||
trace_id=trace_id_hex,
|
||||
span_id=span_id_hex,
|
||||
is_remote=False,
|
||||
trace_state={},
|
||||
)
|
||||
|
||||
# Try to get if span attributes contain something like rollout_id or attempt_id
|
||||
# Override the resource-level attributes with the span-level attributes if present.
|
||||
rollout_id_span = span_attrs.get(SpanNames.ROLLOUT_ID)
|
||||
attempt_id_span = span_attrs.get(SpanNames.ATTEMPT_ID)
|
||||
sequence_id_span = span_attrs.get(SpanNames.SPAN_SEQUENCE_ID)
|
||||
|
||||
# Normalize to regular strings and ints
|
||||
rollout_id_raw = rollout_id_span if rollout_id_span is not None else rollout_id_resource
|
||||
attempt_id_raw = attempt_id_span if attempt_id_span is not None else attempt_id_resource
|
||||
sequence_id_raw = sequence_id_span if sequence_id_span is not None else sequence_id_resource
|
||||
|
||||
rollout_id, attempt_id = _normalize_rollout_attempt_id(rollout_id_raw, attempt_id_raw)
|
||||
sequence_id = _normalize_sequence_id(sequence_id_raw)
|
||||
|
||||
if rollout_id is None or attempt_id is None:
|
||||
logger.warning(
|
||||
"Both rollout_id and attempt_id must be present in resource attributes. "
|
||||
"Spans will not be able to log to the store because of missing IDs: rollout_id=%s, attempt_id=%s, sequence_id=%s",
|
||||
rollout_id,
|
||||
attempt_id,
|
||||
sequence_id,
|
||||
)
|
||||
continue
|
||||
|
||||
# Generate a new sequence ID if not provided
|
||||
if sequence_id is None:
|
||||
current_sequence_id = await store.get_next_span_sequence_id(
|
||||
rollout_id=rollout_id, attempt_id=attempt_id
|
||||
)
|
||||
else:
|
||||
current_sequence_id = sequence_id
|
||||
|
||||
# Build Span
|
||||
span = Span(
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
sequence_id=current_sequence_id,
|
||||
trace_id=trace_id_hex,
|
||||
span_id=span_id_hex,
|
||||
parent_id=parent_id_hex,
|
||||
name=proto_span.name,
|
||||
status=status,
|
||||
attributes=span_attrs,
|
||||
events=_events_from_proto(proto_span),
|
||||
links=_links_from_proto(proto_span),
|
||||
start_time=convert_timestamp(proto_span.start_time_unix_nano),
|
||||
end_time=convert_timestamp(proto_span.end_time_unix_nano),
|
||||
context=context,
|
||||
parent=None, # OTLP only has parent_span_id; we don't have full SpanContext
|
||||
resource=otel_resource,
|
||||
)
|
||||
|
||||
output_spans.append(span)
|
||||
|
||||
return output_spans
|
||||
|
||||
|
||||
class LightningStoreOTLPExporter(OTLPSpanExporter):
|
||||
"""OTLP Exporter that write to a LightningStore-compatible backend.
|
||||
|
||||
The backend requires two special attributes on each span:
|
||||
|
||||
- `agentlightning.rollout_id`: The rollout ID to associate the span with.
|
||||
- `agentlightning.attempt_id`: The attempt ID to associate the span with.
|
||||
|
||||
It can optionally use the following attribute to sequence spans:
|
||||
|
||||
- `agentlightning.span_sequence_id`: A decimal string representing the sequence ID of the span.
|
||||
"""
|
||||
|
||||
_default_endpoint: Optional[str] = None
|
||||
_rollout_id: Optional[str] = None
|
||||
_attempt_id: Optional[str] = None
|
||||
|
||||
def enable_store_otlp(self, endpoint: str, rollout_id: str, attempt_id: str) -> None:
|
||||
"""Enable storing OTLP data to a specific LightningStore rollout/attempt."""
|
||||
self._rollout_id = rollout_id
|
||||
self._attempt_id = attempt_id
|
||||
|
||||
self._default_endpoint = self._endpoint
|
||||
self._endpoint = endpoint
|
||||
|
||||
def disable_store_otlp(self) -> None:
|
||||
"""Disable storing OTLP data to LightningStore."""
|
||||
self._rollout_id = None
|
||||
self._attempt_id = None
|
||||
if self._default_endpoint is not None:
|
||||
self._endpoint = self._default_endpoint
|
||||
|
||||
def should_bypass(self) -> bool:
|
||||
"""Check if the exporter should bypass the default export if rollout_id and attempt_id are not set."""
|
||||
return True
|
||||
|
||||
def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
|
||||
if self._rollout_id is not None and self._attempt_id is not None:
|
||||
# rollout_id and attempt_id are present in resource attributes
|
||||
# It means that the server supports OTLP endpoint.
|
||||
for span in spans:
|
||||
# Override the resources so that the server knows where the request comes from.
|
||||
span._resource = span._resource.merge( # pyright: ignore[reportPrivateUsage]
|
||||
Resource.create(
|
||||
{
|
||||
SpanNames.ROLLOUT_ID: self._rollout_id,
|
||||
SpanNames.ATTEMPT_ID: self._attempt_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
return super().export(spans)
|
||||
elif not self.should_bypass():
|
||||
logger.debug("Rollout ID and Attempt ID not set; using default OTLP exporter behavior.")
|
||||
return super().export(spans)
|
||||
else:
|
||||
logger.debug("Rollout ID and Attempt ID not set; bypassing export.")
|
||||
return SpanExportResult.SUCCESS
|
||||
|
||||
|
||||
def _read_body_maybe_gzip(request: Request, raw_body: bytes) -> bytes:
|
||||
"""
|
||||
Decompress body if Content-Encoding: gzip; otherwise return as is.
|
||||
"""
|
||||
encoding = request.headers.get("Content-Encoding", "").lower()
|
||||
if encoding == "gzip":
|
||||
return gzip.decompress(raw_body)
|
||||
return raw_body
|
||||
|
||||
|
||||
def _maybe_gzip_response(request: Request, payload: bytes) -> Tuple[bytes, Dict[str, str]]:
|
||||
"""
|
||||
If Accept-Encoding includes gzip, gzip the payload and set Content-Encoding header.
|
||||
"""
|
||||
ae = request.headers.get("Accept-Encoding", "")
|
||||
tokens = [token.split(";")[0].strip().lower() for token in ae.split(",") if token.strip()]
|
||||
headers: Dict[str, str] = {}
|
||||
if "gzip" in tokens:
|
||||
payload = gzip.compress(payload)
|
||||
headers["Content-Encoding"] = "gzip"
|
||||
return payload, headers
|
||||
|
||||
|
||||
def _bad_request_response(request: Request, message: str, content_type: str = PROTOBUF_CT) -> Response:
|
||||
"""
|
||||
Build a 400 response whose body is a protobuf Status message, encoded
|
||||
in the same Content-Type as the request (OTLP/HTTP requirement).
|
||||
"""
|
||||
status_msg = Status(message=message)
|
||||
|
||||
if content_type == PROTOBUF_CT:
|
||||
body = status_msg.SerializeToString()
|
||||
else:
|
||||
# Fallback: JSON representation of Status.
|
||||
body = json_format.MessageToJson(status_msg).encode("utf-8")
|
||||
|
||||
body, headers = _maybe_gzip_response(request, body)
|
||||
|
||||
return Response(
|
||||
content=body,
|
||||
status_code=400,
|
||||
media_type=content_type,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_rollout_attempt_id(
|
||||
rollout_id: Optional[AttributeValue], attempt_id: Optional[AttributeValue]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Normalize a rollout or attempt ID to a string."""
|
||||
rollout_id_str = str(rollout_id) if rollout_id is not None else None
|
||||
attempt_id_str = str(attempt_id) if attempt_id is not None else None
|
||||
return rollout_id_str, attempt_id_str
|
||||
|
||||
|
||||
def _normalize_sequence_id(sequence_id: Optional[AttributeValue]) -> Optional[int]:
|
||||
"""Normalize a sequence ID to an integer."""
|
||||
if sequence_id is None:
|
||||
return None
|
||||
try:
|
||||
sequence_id_int = int(str(sequence_id))
|
||||
except (ValueError, TypeError):
|
||||
logger.warning(
|
||||
"Invalid sequence_id value in resource attributes: %r. Must be an integer or string representing an integer. Assuming None.",
|
||||
sequence_id,
|
||||
)
|
||||
sequence_id_int = None
|
||||
return sequence_id_int
|
||||
|
||||
|
||||
def _any_value_to_python(value: AnyValue) -> Any:
|
||||
"""Convert OTLP AnyValue -> plain Python value."""
|
||||
kind = value.WhichOneof("value")
|
||||
if kind is None:
|
||||
return None
|
||||
if kind == "string_value":
|
||||
return value.string_value
|
||||
if kind == "bool_value":
|
||||
return value.bool_value
|
||||
if kind == "int_value":
|
||||
return int(value.int_value)
|
||||
if kind == "double_value":
|
||||
return float(value.double_value)
|
||||
if kind == "array_value":
|
||||
return [_any_value_to_python(v) for v in value.array_value.values]
|
||||
if kind == "kvlist_value":
|
||||
# Map<string, AnyValue> -> dict
|
||||
return {kv.key: _any_value_to_python(kv.value) for kv in value.kvlist_value.values}
|
||||
if kind == "bytes_value":
|
||||
# Serialize bytes as hex string to stay JSON-friendly
|
||||
return value.bytes_value.hex()
|
||||
return None
|
||||
|
||||
|
||||
def _kv_list_to_dict(kvs: Sequence[KeyValue]) -> Attributes:
|
||||
"""Convert repeated KeyValue -> Attributes dict."""
|
||||
return {kv.key: _any_value_to_python(kv.value) for kv in kvs}
|
||||
|
||||
|
||||
_STATUS_CODE_MAP = {
|
||||
ProtoStatus.STATUS_CODE_UNSET: "UNSET",
|
||||
ProtoStatus.STATUS_CODE_OK: "OK",
|
||||
ProtoStatus.STATUS_CODE_ERROR: "ERROR",
|
||||
}
|
||||
|
||||
|
||||
def _bytes_to_trace_id_hex(b: bytes) -> str:
|
||||
# OTLP uses 16-byte trace IDs; format as 32-char hex
|
||||
if not b:
|
||||
return "0" * 32
|
||||
return b.hex().rjust(32, "0")
|
||||
|
||||
|
||||
def _bytes_to_span_id_hex(b: bytes) -> str:
|
||||
# OTLP uses 8-byte span IDs; format as 16-char hex
|
||||
if not b:
|
||||
return "0" * 16
|
||||
return b.hex().rjust(16, "0")
|
||||
|
||||
|
||||
def _events_from_proto(span: ProtoSpan) -> List[Event]:
|
||||
"""Event converter from OTLP ProtoSpan to List[Event]."""
|
||||
return [
|
||||
Event(
|
||||
name=e.name,
|
||||
attributes=_kv_list_to_dict(e.attributes),
|
||||
timestamp=convert_timestamp(e.time_unix_nano),
|
||||
)
|
||||
for e in span.events
|
||||
]
|
||||
|
||||
|
||||
def _links_from_proto(span: ProtoSpan) -> List[Link]:
|
||||
"""Link converter from OTLP ProtoSpan to List[Link]."""
|
||||
links: List[Link] = []
|
||||
for link in span.links:
|
||||
trace_id_hex = _bytes_to_trace_id_hex(link.trace_id)
|
||||
span_id_hex = _bytes_to_span_id_hex(link.span_id)
|
||||
ctx = SpanContext(
|
||||
trace_id=trace_id_hex,
|
||||
span_id=span_id_hex,
|
||||
is_remote=False,
|
||||
trace_state={}, # OTLP trace_state is currently a string; you can parse if needed
|
||||
)
|
||||
links.append(
|
||||
Link(
|
||||
context=ctx,
|
||||
attributes=_kv_list_to_dict(link.attributes) or None,
|
||||
)
|
||||
)
|
||||
return links
|
||||
|
||||
|
||||
def _resource_from_proto(resource: ProtoResource, schema_url: str = "") -> OtelResource:
|
||||
return OtelResource(
|
||||
attributes=_kv_list_to_dict(resource.attributes),
|
||||
schema_url=schema_url or "",
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,72 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import platform
|
||||
import socket
|
||||
from contextlib import suppress
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
import psutil
|
||||
from gpustat import GPUStat, GPUStatCollection
|
||||
|
||||
|
||||
def system_snapshot(include_gpu: bool = False) -> Dict[str, Any]:
|
||||
# CPU
|
||||
cpu = {
|
||||
"cpu_name": platform.processor(),
|
||||
"cpu_cores": psutil.cpu_count(logical=False),
|
||||
"cpu_threads": psutil.cpu_count(logical=True),
|
||||
"cpu_usage_pct": psutil.cpu_percent(0.05),
|
||||
}
|
||||
|
||||
# Memory
|
||||
vm = psutil.virtual_memory()
|
||||
mem = {
|
||||
"mem_used_gb": round(vm.used / (2**30), 2),
|
||||
"mem_total_gb": round(vm.total / (2**30), 2),
|
||||
"mem_pct": vm.percent,
|
||||
}
|
||||
|
||||
# Disk
|
||||
du = psutil.disk_usage("/")
|
||||
disk = {
|
||||
"disk_used_gb": round(du.used / (2**30), 2),
|
||||
"disk_total_gb": round(du.total / (2**30), 2),
|
||||
"disk_pct": du.percent,
|
||||
}
|
||||
|
||||
# GPU
|
||||
gpus: List[Dict[str, Any]] = []
|
||||
with suppress(Exception):
|
||||
for g in GPUStatCollection.new_query().gpus: # type: ignore
|
||||
g = cast(GPUStat, g)
|
||||
gpus.append(
|
||||
{
|
||||
"gpu": g.name, # type: ignore
|
||||
"util_pct": g.utilization,
|
||||
"mem_used_mb": g.memory_used,
|
||||
"mem_total_mb": g.memory_total,
|
||||
"temp_c": g.temperature,
|
||||
}
|
||||
)
|
||||
|
||||
# Network
|
||||
net = psutil.net_io_counters()
|
||||
netinfo = {
|
||||
"bytes_sent_mb": round(net.bytes_sent / (2**20), 2),
|
||||
"bytes_recv_mb": round(net.bytes_recv / (2**20), 2),
|
||||
}
|
||||
|
||||
# OS / meta
|
||||
return {
|
||||
"timestamp": datetime.now().isoformat(timespec="seconds"),
|
||||
"host": socket.gethostname(),
|
||||
"os": platform.platform(),
|
||||
**cpu,
|
||||
**mem,
|
||||
**disk,
|
||||
**netinfo,
|
||||
**({"gpus": gpus} if include_gpu else {}),
|
||||
}
|
||||
@@ -18,13 +18,13 @@ from flask import Flask, Response, abort, request
|
||||
from tensordict import TensorDict
|
||||
from verl import DataProto
|
||||
|
||||
from agentlightning import LLM, AgentLightningServer, NamedResources, RolloutLegacy, configure_logger
|
||||
from agentlightning import LLM, AgentLightningServer, NamedResources, RolloutLegacy, setup_logging
|
||||
from agentlightning.adapter.triplet import TracerTraceToTriplet, TraceToTripletBase
|
||||
from agentlightning.llm_proxy import LLMProxy, ModelConfig
|
||||
from agentlightning.store.base import LightningStore
|
||||
from agentlightning.types import Rollout, RolloutConfig, Task
|
||||
|
||||
configure_logger()
|
||||
setup_logging()
|
||||
|
||||
__all__ = [
|
||||
"AgentModeDaemon",
|
||||
@@ -294,7 +294,7 @@ class AgentModeDaemon:
|
||||
self._proxy_thread.start()
|
||||
print(f"Proxy server running on port {self.proxy_port}")
|
||||
|
||||
def _update_proxy_server_v1(self):
|
||||
async def _update_proxy_server_v1(self):
|
||||
model_name = self.train_information.get("model")
|
||||
if not model_name:
|
||||
raise ValueError("Model name is not set.")
|
||||
@@ -313,12 +313,7 @@ class AgentModeDaemon:
|
||||
],
|
||||
)
|
||||
|
||||
if self.llm_proxy.is_running():
|
||||
# FIXME: Need to switch to a different port right now
|
||||
# because the forked processes carried the old fd
|
||||
self.llm_proxy.restart(_port=_find_available_port())
|
||||
else:
|
||||
self.llm_proxy.start()
|
||||
await self.llm_proxy.restart()
|
||||
|
||||
def start(self):
|
||||
"""Starts the main AgentLightningServer and the proxy server."""
|
||||
@@ -352,7 +347,7 @@ class AgentModeDaemon:
|
||||
if server_addresses != self.backend_llm_server_addresses:
|
||||
self.backend_llm_server_addresses = server_addresses
|
||||
if self.mode == "v1" and not self.llm_proxy.is_running():
|
||||
self._update_proxy_server_v1()
|
||||
await self._update_proxy_server_v1()
|
||||
self.is_train = is_train
|
||||
|
||||
# 1. Update resources on the server for clients to use
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
|
||||
# type: ignore
|
||||
|
||||
from importlib.metadata import version
|
||||
from typing import Any
|
||||
|
||||
import hydra
|
||||
import ray
|
||||
from packaging import version as packaging_version
|
||||
from verl.trainer.main_ppo import create_rl_sampler
|
||||
from verl.trainer.ppo.reward import load_reward_manager
|
||||
|
||||
@@ -39,11 +41,17 @@ def run_ppo(
|
||||
) -> None:
|
||||
if not ray.is_initialized():
|
||||
# this is for local ray cluster
|
||||
try:
|
||||
# verl >= 0.6.0
|
||||
num_cpus = config.ray_kwargs.ray_init.num_cpus
|
||||
except AttributeError:
|
||||
# verl < 0.6.0
|
||||
num_cpus = config.ray_init.num_cpus
|
||||
ray.init(
|
||||
runtime_env={
|
||||
"env_vars": {"TOKENIZERS_PARALLELISM": "true", "NCCL_DEBUG": "WARN", "VLLM_LOGGING_LEVEL": "WARN"}
|
||||
},
|
||||
num_cpus=config.ray_init.num_cpus,
|
||||
num_cpus=num_cpus,
|
||||
)
|
||||
|
||||
runner = TaskRunner.remote()
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Dict, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import verl
|
||||
from codetiming import Timer
|
||||
from omegaconf import OmegaConf
|
||||
from tqdm import tqdm
|
||||
@@ -19,7 +20,7 @@ from verl import DataProto
|
||||
from verl.protocol import pad_dataproto_to_divisor, unpad_dataproto
|
||||
from verl.trainer.ppo.core_algos import agg_loss
|
||||
from verl.trainer.ppo.metric_utils import (
|
||||
compute_data_metrics,
|
||||
_compute_response_info,
|
||||
compute_throughout_metrics,
|
||||
compute_timing_metrics,
|
||||
)
|
||||
@@ -53,6 +54,108 @@ def _timer(name: str, timing_raw: Dict[str, float]):
|
||||
timing_raw[name] += timer.last
|
||||
|
||||
|
||||
# This function is adapted from verl.
|
||||
# We introduce a new parameter `suffix` to distinguish between metrics computed
|
||||
# before and after AgentLightning’s post-processing.
|
||||
# - "Before" refers to raw reward and advantage values.
|
||||
# - "After" refers to values computed following post-processing, which involves:
|
||||
# (1) Dropping prompts that exceed the maximum allowed length.
|
||||
# (2) Adjusting the batch size to be a multiple of the mini PPO size.
|
||||
# Different suffixes are used to label these two stages accordingly.
|
||||
def compute_data_metrics(batch: DataProto, use_critic: bool = True, suffix: str = "") -> Dict[str, Any]:
|
||||
"""
|
||||
Computes various metrics from a batch of data for PPO training.
|
||||
|
||||
This function calculates metrics related to scores, rewards, advantages, returns, values,
|
||||
and sequence lengths from a batch of data. It provides statistical information (mean, max, min)
|
||||
for each metric category.
|
||||
|
||||
Args:
|
||||
batch: A DataProto object containing batch data with token-level scores, rewards, advantages, etc.
|
||||
use_critic: Whether to include critic-specific metrics. Defaults to True.
|
||||
|
||||
Returns:
|
||||
A dictionary of metrics including:
|
||||
- critic/score/mean, max, min: Statistics about sequence scores
|
||||
- critic/rewards/mean, max, min: Statistics about sequence rewards
|
||||
- critic/advantages/mean, max, min: Statistics about advantages
|
||||
- critic/returns/mean, max, min: Statistics about returns
|
||||
- critic/values/mean, max, min: Statistics about critic values (if use_critic=True)
|
||||
- critic/vf_explained_var: Explained variance of the value function (if use_critic=True)
|
||||
- response_length/mean, max, min, clip_ratio: Statistics about response lengths
|
||||
- prompt_length/mean, max, min, clip_ratio: Statistics about prompt lengths
|
||||
"""
|
||||
sequence_score = batch.batch["token_level_scores"].sum(-1)
|
||||
sequence_reward = batch.batch["token_level_rewards"].sum(-1)
|
||||
|
||||
advantages = batch.batch["advantages"]
|
||||
returns = batch.batch["returns"]
|
||||
|
||||
max_response_length = batch.batch["responses"].shape[-1]
|
||||
|
||||
prompt_mask = batch.batch["attention_mask"][:, :-max_response_length].bool()
|
||||
response_mask = batch.batch["attention_mask"][:, -max_response_length:].bool()
|
||||
|
||||
max_prompt_length = prompt_mask.size(-1)
|
||||
|
||||
response_info = _compute_response_info(batch)
|
||||
prompt_length = response_info["prompt_length"]
|
||||
response_length = response_info["response_length"]
|
||||
|
||||
valid_adv = torch.masked_select(advantages, response_mask)
|
||||
valid_returns = torch.masked_select(returns, response_mask)
|
||||
|
||||
if use_critic:
|
||||
values = batch.batch["values"]
|
||||
valid_values = torch.masked_select(values, response_mask)
|
||||
return_diff_var = torch.var(valid_returns - valid_values)
|
||||
return_var = torch.var(valid_returns)
|
||||
|
||||
metrics = {
|
||||
# score
|
||||
"critic/score/mean" + suffix: torch.mean(sequence_score).detach().item(),
|
||||
"critic/score/max" + suffix: torch.max(sequence_score).detach().item(),
|
||||
"critic/score/min" + suffix: torch.min(sequence_score).detach().item(),
|
||||
# reward
|
||||
"critic/rewards/mean" + suffix: torch.mean(sequence_reward).detach().item(),
|
||||
"critic/rewards/max" + suffix: torch.max(sequence_reward).detach().item(),
|
||||
"critic/rewards/min" + suffix: torch.min(sequence_reward).detach().item(),
|
||||
# adv
|
||||
"critic/advantages/mean" + suffix: torch.mean(valid_adv).detach().item(),
|
||||
"critic/advantages/max" + suffix: torch.max(valid_adv).detach().item(),
|
||||
"critic/advantages/min" + suffix: torch.min(valid_adv).detach().item(),
|
||||
# returns
|
||||
"critic/returns/mean" + suffix: torch.mean(valid_returns).detach().item(),
|
||||
"critic/returns/max" + suffix: torch.max(valid_returns).detach().item(),
|
||||
"critic/returns/min" + suffix: torch.min(valid_returns).detach().item(),
|
||||
**(
|
||||
{
|
||||
# values
|
||||
"critic/values/mean" + suffix: torch.mean(valid_values).detach().item(),
|
||||
"critic/values/max" + suffix: torch.max(valid_values).detach().item(),
|
||||
"critic/values/min" + suffix: torch.min(valid_values).detach().item(),
|
||||
# vf explained var
|
||||
"critic/vf_explained_var" + suffix: (1.0 - return_diff_var / (return_var + 1e-5)).detach().item(),
|
||||
}
|
||||
if use_critic
|
||||
else {}
|
||||
),
|
||||
# response length
|
||||
"response_length/mean" + suffix: torch.mean(response_length).detach().item(),
|
||||
"response_length/max" + suffix: torch.max(response_length).detach().item(),
|
||||
"response_length/min" + suffix: torch.min(response_length).detach().item(),
|
||||
"response_length/clip_ratio"
|
||||
+ suffix: torch.mean(torch.eq(response_length, max_response_length).float()).detach().item(),
|
||||
# prompt length
|
||||
"prompt_length/mean" + suffix: torch.mean(prompt_length).detach().item(),
|
||||
"prompt_length/max" + suffix: torch.max(prompt_length).detach().item(),
|
||||
"prompt_length/min" + suffix: torch.min(prompt_length).detach().item(),
|
||||
"prompt_length/clip_ratio"
|
||||
+ suffix: torch.mean(torch.eq(prompt_length, max_prompt_length).float()).detach().item(),
|
||||
}
|
||||
return metrics
|
||||
|
||||
|
||||
class AgentLightningTrainer(RayPPOTrainer):
|
||||
"""
|
||||
Specialized PPO trainer for agent-based reinforcement learning.
|
||||
@@ -215,6 +318,9 @@ class AgentLightningTrainer(RayPPOTrainer):
|
||||
config=self.config.algorithm,
|
||||
)
|
||||
|
||||
# Calculate the metrics before processing. Refer to the comments of function `compute_data_metrics` for details.
|
||||
metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic, suffix="_before_processing"))
|
||||
|
||||
# after advantages are assinged, we begin to drop (1) long prompt (2) floor to ppo minisize
|
||||
keep_indices = (~batch.batch["is_drop_mask"]).nonzero(as_tuple=True)[0]
|
||||
metrics["training/n_triplets_prompt_too_long"] = (
|
||||
@@ -274,7 +380,7 @@ class AgentLightningTrainer(RayPPOTrainer):
|
||||
)
|
||||
|
||||
# compute training metrics
|
||||
metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic))
|
||||
metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic, suffix="_after_processing"))
|
||||
metrics.update(compute_timing_metrics(batch=batch, timing_raw=timing_raw))
|
||||
# TODO: implement actual tflpo and theoretical tflpo
|
||||
n_gpus = self.resource_pool_manager.get_n_gpus()
|
||||
@@ -298,14 +404,20 @@ class AgentLightningTrainer(RayPPOTrainer):
|
||||
assert self.async_rollout_mode, "If agent mode is enabled, async server must be enabled"
|
||||
if self.adapter is not None and not isinstance(self.adapter, TraceToTripletBase):
|
||||
raise ValueError("Adapter must be a TraceToTripletBase for currently VERL implementation.")
|
||||
verl_version = verl.__version__
|
||||
if verl_version == "0.5.0":
|
||||
# Note (Zhiyuan): To avoid further patch into vllm async server, using the same sentence to get the naming here.
|
||||
# However, it is possible that verl updates the naming and causes incompatibility.
|
||||
# Reference: https://github.com/volcengine/verl/blob/5b5e09d9cc20625e436d01f69d9cc739ff681c54/verl/workers/rollout/vllm_rollout/vllm_async_server.py#L217
|
||||
model = "/".join(self.config.actor_rollout_ref.model.path.split("/")[-2:])
|
||||
else:
|
||||
# For other versions (e.g., 0.6.0), we use the full path to the model.
|
||||
model = self.config.actor_rollout_ref.model.path
|
||||
self.agent_mode_daemon = AgentModeDaemon(
|
||||
self.config.agentlightning.port,
|
||||
self.config.actor_rollout_ref.rollout.n,
|
||||
train_information={
|
||||
# Note (Zhiyuan): To avoid further patch into vllm async server, using the same sentence to get the naming here.
|
||||
# However, it is possible that verl updates the naming and causes incompatibility.
|
||||
# Reference: https://github.com/volcengine/verl/blob/5b5e09d9cc20625e436d01f69d9cc739ff681c54/verl/workers/rollout/vllm_rollout/vllm_async_server.py#L217
|
||||
"model": "/".join(self.config.actor_rollout_ref.model.path.split("/")[-2:]),
|
||||
"model": model,
|
||||
"temperature": self.config.actor_rollout_ref.rollout.temperature,
|
||||
},
|
||||
tokenizer=self.tokenizer,
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
# Logs
|
||||
logs
|
||||
*.log
|
||||
npm-debug.log*
|
||||
yarn-debug.log*
|
||||
yarn-error.log*
|
||||
lerna-debug.log*
|
||||
.pnpm-debug.log*
|
||||
|
||||
# Diagnostic reports (https://nodejs.org/api/report.html)
|
||||
report.[0-9]*.[0-9]*.[0-9]*.[0-9]*.json
|
||||
|
||||
# Runtime data
|
||||
pids
|
||||
*.pid
|
||||
*.seed
|
||||
*.pid.lock
|
||||
|
||||
# Directory for instrumented libs generated by jscoverage/JSCover
|
||||
lib-cov
|
||||
|
||||
# Coverage directory used by tools like istanbul
|
||||
coverage
|
||||
*.lcov
|
||||
|
||||
# nyc test coverage
|
||||
.nyc_output
|
||||
|
||||
# Grunt intermediate storage (https://gruntjs.com/creating-plugins#storing-task-files)
|
||||
.grunt
|
||||
|
||||
# Bower dependency directory (https://bower.io/)
|
||||
bower_components
|
||||
|
||||
# node-waf configuration
|
||||
.lock-wscript
|
||||
|
||||
# Compiled binary addons (https://nodejs.org/api/addons.html)
|
||||
build/Release
|
||||
|
||||
# Dependency directories
|
||||
node_modules/
|
||||
jspm_packages/
|
||||
|
||||
# Snowpack dependency directory (https://snowpack.dev/)
|
||||
web_modules/
|
||||
|
||||
# TypeScript cache
|
||||
*.tsbuildinfo
|
||||
|
||||
# Optional npm cache directory
|
||||
.npm
|
||||
|
||||
# Optional eslint cache
|
||||
.eslintcache
|
||||
|
||||
# Optional stylelint cache
|
||||
.stylelintcache
|
||||
|
||||
# Microbundle cache
|
||||
.rpt2_cache/
|
||||
.rts2_cache_cjs/
|
||||
.rts2_cache_es/
|
||||
.rts2_cache_umd/
|
||||
|
||||
# Optional REPL history
|
||||
.node_repl_history
|
||||
|
||||
# Output of 'npm pack'
|
||||
*.tgz
|
||||
|
||||
# Yarn Integrity file
|
||||
.yarn-integrity
|
||||
|
||||
# dotenv environment variable files
|
||||
.env
|
||||
.env.development.local
|
||||
.env.test.local
|
||||
.env.production.local
|
||||
.env.local
|
||||
|
||||
# parcel-bundler cache (https://parceljs.org/)
|
||||
.cache
|
||||
.parcel-cache
|
||||
|
||||
# Next.js build output
|
||||
.next
|
||||
out
|
||||
|
||||
# Nuxt.js build / generate output
|
||||
.nuxt
|
||||
dist
|
||||
|
||||
# Gatsby files
|
||||
.cache/
|
||||
# Comment in the public line in if your project uses Gatsby and not Next.js
|
||||
# https://nextjs.org/blog/next-9-1#public-directory-support
|
||||
# public
|
||||
|
||||
# vuepress build output
|
||||
.vuepress/dist
|
||||
|
||||
# vuepress v2.x temp and cache directory
|
||||
.temp
|
||||
.cache
|
||||
|
||||
# Docusaurus cache and generated files
|
||||
.docusaurus
|
||||
|
||||
# Serverless directories
|
||||
.serverless/
|
||||
|
||||
# FuseBox cache
|
||||
.fusebox/
|
||||
|
||||
# DynamoDB Local files
|
||||
.dynamodb/
|
||||
|
||||
# TernJS port file
|
||||
.tern-port
|
||||
|
||||
# Stores VSCode versions used for testing VSCode extensions
|
||||
.vscode-test
|
||||
|
||||
# yarn v2
|
||||
.yarn/cache
|
||||
.yarn/unplugged
|
||||
.yarn/build-state.yml
|
||||
.yarn/install-state.gz
|
||||
.pnp.*
|
||||
|
||||
.DS_Store
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
/** @type {import("@ianvs/prettier-plugin-sort-imports").PrettierConfig} */
|
||||
const config = {
|
||||
printWidth: 120,
|
||||
singleQuote: true,
|
||||
tabWidth: 2,
|
||||
useTabs: false,
|
||||
semi: true,
|
||||
quoteProps: 'consistent',
|
||||
jsxSingleQuote: true,
|
||||
trailingComma: 'all',
|
||||
bracketSpacing: true,
|
||||
objectWrap: 'preserve',
|
||||
arrowParens: 'always',
|
||||
proseWrap: 'preserve',
|
||||
endOfLine: 'lf',
|
||||
plugins: ['@ianvs/prettier-plugin-sort-imports'],
|
||||
importOrder: [
|
||||
'.*styles.css$',
|
||||
'',
|
||||
'dayjs',
|
||||
'^react$',
|
||||
'^next$',
|
||||
'^next/.*$',
|
||||
'<BUILTIN_MODULES>',
|
||||
'<THIRD_PARTY_MODULES>',
|
||||
'^@mantine/(.*)$',
|
||||
'^@mantinex/(.*)$',
|
||||
'^@mantine-tests/(.*)$',
|
||||
'^@docs/(.*)$',
|
||||
'^@/.*$',
|
||||
'^../(?!.*.css$).*$',
|
||||
'^./(?!.*.css$).*$',
|
||||
'\\.css$',
|
||||
],
|
||||
overrides: [
|
||||
{
|
||||
files: '*.mdx',
|
||||
options: {
|
||||
printWidth: 120,
|
||||
},
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
export default config;
|
||||
@@ -0,0 +1,12 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
// Centralized constants that keep Storybook fixtures deterministic so Chromatic
|
||||
// snapshots do not drift when the build environment changes.
|
||||
export const STORY_DATE_NOW_MS = 1762775145209;
|
||||
export const STORY_DATE_NOW_SECONDS = Math.floor(STORY_DATE_NOW_MS / 1000);
|
||||
|
||||
// Use a fixed origin so any code that would normally read window.location.*
|
||||
// in the app can rely on the same value from Storybook fixtures. Prefer HTTPS
|
||||
// so Chromatic (which is served over HTTPS) avoids mixed-content fetch errors.
|
||||
export const STORY_BASE_URL = 'https://storybook.agentlightning.invalid';
|
||||
export const STORY_LOCATION_HREF = `${STORY_BASE_URL}/storybook`;
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { StorybookConfig } from '@storybook/react-vite';
|
||||
|
||||
const config: StorybookConfig = {
|
||||
core: {
|
||||
disableWhatsNewNotifications: true,
|
||||
disableTelemetry: true,
|
||||
enableCrashReports: false,
|
||||
},
|
||||
stories: ['../src/**/*.mdx', '../src/**/*.story.@(js|jsx|ts|tsx)'],
|
||||
staticDirs: ['../static'],
|
||||
addons: ['@storybook/addon-themes', '@storybook/addon-vitest'],
|
||||
framework: {
|
||||
name: '@storybook/react-vite',
|
||||
options: {},
|
||||
},
|
||||
};
|
||||
|
||||
export default config;
|
||||
@@ -0,0 +1,13 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
export const allModes = {
|
||||
MD: {
|
||||
viewport: 'md',
|
||||
},
|
||||
LG: {
|
||||
viewport: 'lg',
|
||||
},
|
||||
XL: {
|
||||
viewport: 'xl',
|
||||
},
|
||||
} as const;
|
||||
@@ -0,0 +1,79 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import '@mantine/core/styles.css';
|
||||
import 'mantine-datatable/styles.css';
|
||||
import '../src/styles/theme.css';
|
||||
import '../src/styles/app.css';
|
||||
|
||||
import { initialize, mswLoader } from 'msw-storybook-addon';
|
||||
import { ColorSchemeScript, MantineProvider } from '@mantine/core';
|
||||
import { shadcnCssVariableResolver } from '../src/cssVariableResolver';
|
||||
import { theme as mantineTheme } from '../src/theme';
|
||||
import { STORY_DATE_NOW_MS } from './constants';
|
||||
|
||||
type ColorSchemeValue = 'light' | 'dark';
|
||||
|
||||
initialize({
|
||||
onUnhandledRequest: 'bypass',
|
||||
serviceWorker: {
|
||||
url: '/mockServiceWorker.js',
|
||||
},
|
||||
});
|
||||
|
||||
const fixedDateNow = (() => {
|
||||
const patched = Date.now as typeof Date.now & { __storybookPatched?: boolean };
|
||||
if (patched.__storybookPatched) {
|
||||
return patched;
|
||||
}
|
||||
const replacement = (() => STORY_DATE_NOW_MS) as typeof Date.now & { __storybookPatched?: boolean };
|
||||
replacement.__storybookPatched = true;
|
||||
return replacement;
|
||||
})();
|
||||
|
||||
Date.now = fixedDateNow;
|
||||
|
||||
export const parameters = {
|
||||
layout: 'fullscreen',
|
||||
options: {
|
||||
showPanel: false,
|
||||
// @ts-expect-error – storybook throws build error for (a: any, b: any)
|
||||
storySort: (a, b) => a.title.localeCompare(b.title, undefined, { numeric: true }),
|
||||
},
|
||||
backgrounds: { disable: true },
|
||||
viewport: {
|
||||
options: {
|
||||
md: { name: 'md', styles: { width: '1280px', height: '800px' } },
|
||||
lg: { name: 'lg', styles: { width: '1920px', height: '1080px' } },
|
||||
xl: { name: 'xl', styles: { width: '2560px', height: '1440px' } },
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const globalTypes = {
|
||||
theme: {
|
||||
name: 'Theme',
|
||||
description: 'Mantine color scheme',
|
||||
defaultValue: 'light',
|
||||
toolbar: {
|
||||
icon: 'mirror',
|
||||
items: [
|
||||
{ value: 'light', title: 'Light' },
|
||||
{ value: 'dark', title: 'Dark' },
|
||||
],
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const decorators = [
|
||||
(Story: any, context: any) => {
|
||||
const scheme = (context.parameters.theme ?? context.globals.theme ?? 'light') as ColorSchemeValue;
|
||||
return (
|
||||
<MantineProvider theme={mantineTheme} cssVariablesResolver={shadcnCssVariableResolver} forceColorScheme={scheme}>
|
||||
<ColorSchemeScript />
|
||||
<Story />
|
||||
</MantineProvider>
|
||||
);
|
||||
},
|
||||
];
|
||||
|
||||
export const loaders = [mswLoader];
|
||||
@@ -0,0 +1,8 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { setProjectAnnotations } from '@storybook/react-vite';
|
||||
import * as projectAnnotations from './preview';
|
||||
|
||||
// This is an important step to apply the right configuration when testing your stories.
|
||||
// More info at: https://storybook.js.org/docs/api/portable-stories/portable-stories-vitest#setprojectannotations
|
||||
setProjectAnnotations([projectAnnotations]);
|
||||
@@ -0,0 +1,5 @@
|
||||
# Generated files
|
||||
dist
|
||||
|
||||
# Theme files
|
||||
theme.css
|
||||
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"extends": ["stylelint-config-standard-scss"],
|
||||
"rules": {
|
||||
"custom-property-pattern": null,
|
||||
"selector-class-pattern": null,
|
||||
"scss/no-duplicate-mixins": null,
|
||||
"declaration-empty-line-before": null,
|
||||
"declaration-block-no-redundant-longhand-properties": null,
|
||||
"alpha-value-notation": null,
|
||||
"custom-property-empty-line-before": null,
|
||||
"property-no-vendor-prefix": null,
|
||||
"color-function-notation": null,
|
||||
"length-zero-no-unit": null,
|
||||
"selector-not-notation": null,
|
||||
"no-descending-specificity": null,
|
||||
"comment-empty-line-before": null,
|
||||
"scss/at-mixin-pattern": null,
|
||||
"scss/at-rule-no-unknown": null,
|
||||
"value-keyword-case": null,
|
||||
"media-feature-range-notation": null,
|
||||
"selector-pseudo-class-no-unknown": [
|
||||
true,
|
||||
{
|
||||
"ignorePseudoClasses": ["global"]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
# Agent-lightning Dashboard
|
||||
|
||||
This is the dashboard for Agent-lightning. It is a web application that allows you to inspect your Agent-lightning store and debug running experiments.
|
||||
|
||||
The dashboard is built with React, Mantine UI, and Storybook.
|
||||
|
||||
## npm scripts
|
||||
|
||||
## Build and dev scripts
|
||||
|
||||
- `dev` – start development server
|
||||
- `build` – build production version of the app
|
||||
- `preview` – locally preview production build
|
||||
|
||||
### Testing scripts
|
||||
|
||||
- `eslint` - runs ESLint
|
||||
- `stylelint` - runs Stylelint
|
||||
- `prettier` - runs Prettier
|
||||
- `typecheck` - runs TypeScript typecheck
|
||||
- `vitest` – runs vitest tests
|
||||
- `chromatic` – runs chromatic tests
|
||||
|
||||
### Other scripts
|
||||
|
||||
- `storybook` – starts storybook dev server
|
||||
- `build-storybook` – build production storybook bundle to `storybook-static`
|
||||
@@ -0,0 +1,50 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
// @ts-check
|
||||
import stylistic from '@stylistic/eslint-plugin';
|
||||
import mantine from 'eslint-config-mantine';
|
||||
import { defineConfig } from 'eslint/config';
|
||||
import tseslint from 'typescript-eslint';
|
||||
|
||||
export default defineConfig([
|
||||
// These are arrays → safe to spread
|
||||
...tseslint.configs.recommended,
|
||||
stylistic.configs.customize({ semi: true }),
|
||||
|
||||
// mantine is often a single object → include as-is (or spread only if it's actually an array)
|
||||
...(Array.isArray(mantine) ? mantine : [mantine]),
|
||||
|
||||
// ignores go as their own entry
|
||||
{ ignores: ['**/*.{mjs,cjs,js,d.ts,d.mts}'] },
|
||||
|
||||
// file-specific rules
|
||||
{
|
||||
files: ['**/*.story.tsx'],
|
||||
rules: { 'no-console': 'off' },
|
||||
},
|
||||
|
||||
// project/TS settings + your custom rules
|
||||
{
|
||||
languageOptions: {
|
||||
parserOptions: {
|
||||
tsconfigRootDir: process.cwd(),
|
||||
project: ['./tsconfig.json'],
|
||||
},
|
||||
},
|
||||
rules: {
|
||||
// Disabling conflict rules with prettier
|
||||
'@stylistic/brace-style': ['error', '1tbs', { allowSingleLine: false }],
|
||||
'@stylistic/no-trailing-spaces': 'error',
|
||||
'@stylistic/no-multiple-empty-lines': ['error', { max: 2, maxEOF: 1 }],
|
||||
'@stylistic/jsx-quotes': ['error', 'prefer-single'],
|
||||
'@stylistic/multiline-ternary': 'off',
|
||||
'@stylistic/arrow-parens': ['error', 'always'],
|
||||
'@stylistic/jsx-closing-bracket-location': 'off',
|
||||
'@stylistic/operator-linebreak': 'off',
|
||||
'@stylistic/jsx-newline': 'off',
|
||||
'@stylistic/jsx-one-expression-per-line': 'off',
|
||||
'@stylistic/indent': 'off',
|
||||
'@stylistic/indent-binary-ops': 'off',
|
||||
},
|
||||
},
|
||||
]);
|
||||
Generated
+10890
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,75 @@
|
||||
{
|
||||
"name": "agent-lightning-dashboard",
|
||||
"type": "module",
|
||||
"version": "0.3.0",
|
||||
"scripts": {
|
||||
"dev": "vite",
|
||||
"build": "tsc && vite build",
|
||||
"preview": "vite preview",
|
||||
"typecheck": "tsc --noEmit",
|
||||
"eslint": "eslint .",
|
||||
"stylelint": "stylelint '**/*.css'",
|
||||
"prettier": "prettier --check \"**/*.{ts,tsx,mjs,cjs}\"",
|
||||
"vitest": "vitest run --project unit",
|
||||
"vitest-storybook": "vitest run --project storybook",
|
||||
"storybook": "storybook dev -p 6006",
|
||||
"build-storybook": "storybook build",
|
||||
"chromatic": "chromatic"
|
||||
},
|
||||
"dependencies": {
|
||||
"@mantine/core": "8.3.5",
|
||||
"@mantine/hooks": "8.3.5",
|
||||
"@monaco-editor/react": "^4.7.0",
|
||||
"@reduxjs/toolkit": "^2.9.2",
|
||||
"@tabler/icons-react": "^3.35.0",
|
||||
"clsx": "^2.1.1",
|
||||
"dayjs": "^1.11.18",
|
||||
"mantine-datatable": "^8.2.0",
|
||||
"react": "^19.2.0",
|
||||
"react-dom": "^19.2.0",
|
||||
"react-redux": "^9.2.0",
|
||||
"react-router-dom": "^7.9.4"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@eslint/js": "^9.37.0",
|
||||
"@ianvs/prettier-plugin-sort-imports": "^4.7.0",
|
||||
"@storybook/addon-themes": "^9.1.10",
|
||||
"@storybook/addon-vitest": "^9.1.16",
|
||||
"@storybook/react": "^9.1.10",
|
||||
"@storybook/react-vite": "^9.1.10",
|
||||
"@stylistic/eslint-plugin": "^5.5.0",
|
||||
"@testing-library/dom": "^10.4.1",
|
||||
"@testing-library/jest-dom": "^6.9.1",
|
||||
"@testing-library/react": "^16.3.0",
|
||||
"@testing-library/user-event": "^14.6.1",
|
||||
"@types/node": "^24.7.1",
|
||||
"@types/react": "^19.2.2",
|
||||
"@types/react-dom": "^19.2.1",
|
||||
"@vitejs/plugin-react": "^5.0.4",
|
||||
"chromatic": "^13.3.3",
|
||||
"eslint": "^9.37.0",
|
||||
"eslint-config-mantine": "^4.0.3",
|
||||
"eslint-plugin-jsx-a11y": "^6.10.2",
|
||||
"eslint-plugin-react": "^7.37.5",
|
||||
"identity-obj-proxy": "^3.0.0",
|
||||
"jsdom": "^27.0.0",
|
||||
"msw": "^2.11.6",
|
||||
"msw-storybook-addon": "^2.0.6",
|
||||
"postcss": "^8.5.6",
|
||||
"postcss-preset-mantine": "1.18.0",
|
||||
"postcss-simple-vars": "^7.0.1",
|
||||
"prettier": "^3.6.2",
|
||||
"prop-types": "^15.8.1",
|
||||
"storybook": "^9.1.10",
|
||||
"stylelint": "^16.25.0",
|
||||
"stylelint-config-standard-scss": "^16.0.0",
|
||||
"typescript": "^5.9.3",
|
||||
"typescript-eslint": "^8.46.0",
|
||||
"vite": "^7.1.9",
|
||||
"vite-tsconfig-paths": "^5.1.4",
|
||||
"vitest": "^4.0.0",
|
||||
"playwright": "^1.56.1",
|
||||
"@vitest/browser-playwright": "4.0.4",
|
||||
"@vitest/coverage-v8": "4.0.4"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
module.exports = {
|
||||
plugins: {
|
||||
'postcss-preset-mantine': {},
|
||||
'postcss-simple-vars': {
|
||||
variables: {
|
||||
'mantine-breakpoint-xs': '36em',
|
||||
'mantine-breakpoint-sm': '48em',
|
||||
'mantine-breakpoint-md': '62em',
|
||||
'mantine-breakpoint-lg': '75em',
|
||||
'mantine-breakpoint-xl': '88em',
|
||||
},
|
||||
},
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,13 @@
|
||||
<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<link rel="icon" type="image/svg+xml" href="../src/favicon.svg" />
|
||||
<meta name="viewport" content="minimum-scale=1, initial-scale=1, width=device-width, user-scalable=no" />
|
||||
<title>Agent-lightning Dashboard</title>
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
<script type="module" src="/main.tsx"></script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,3 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import '../src/main.js';
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import '@mantine/core/styles.css';
|
||||
import 'mantine-datatable/styles.css';
|
||||
import './styles/theme.css';
|
||||
import './styles/app.css';
|
||||
|
||||
import { MantineProvider } from '@mantine/core';
|
||||
import { useColorScheme } from '@mantine/hooks';
|
||||
import { shadcnCssVariableResolver } from './cssVariableResolver';
|
||||
import { selectThemePreference } from './features/config/selectors';
|
||||
import { Router } from './Router';
|
||||
import { useAppSelector } from './store/hooks';
|
||||
import { shadcnTheme } from './theme';
|
||||
|
||||
export default function App() {
|
||||
const themePreference = useAppSelector(selectThemePreference);
|
||||
const systemColorScheme = useColorScheme();
|
||||
const resolvedColorScheme = themePreference === 'system' ? systemColorScheme : themePreference;
|
||||
|
||||
return (
|
||||
<MantineProvider
|
||||
theme={shadcnTheme}
|
||||
cssVariablesResolver={shadcnCssVariableResolver}
|
||||
forceColorScheme={resolvedColorScheme}
|
||||
>
|
||||
<Router />
|
||||
</MantineProvider>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createBrowserRouter, Navigate, RouterProvider } from 'react-router-dom';
|
||||
import { AppLayoutWithState } from './layouts/AppLayout';
|
||||
import { ResourcesPage } from './pages/Resources.page';
|
||||
import { RolloutsPage } from './pages/Rollouts.page';
|
||||
import { SettingsPage } from './pages/Settings.page';
|
||||
import { TracesPage } from './pages/Traces.page';
|
||||
import { WorkersPage } from './pages/Workers.page';
|
||||
|
||||
const router = createBrowserRouter([
|
||||
{
|
||||
path: '/',
|
||||
element: <AppLayoutWithState />,
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
element: <Navigate to='/rollouts' replace />,
|
||||
},
|
||||
{
|
||||
path: 'rollouts',
|
||||
element: <RolloutsPage />,
|
||||
},
|
||||
{
|
||||
path: 'resources',
|
||||
element: <ResourcesPage />,
|
||||
},
|
||||
{
|
||||
path: 'traces',
|
||||
element: <TracesPage />,
|
||||
},
|
||||
{
|
||||
path: 'runners',
|
||||
element: <WorkersPage />,
|
||||
},
|
||||
{
|
||||
path: 'settings',
|
||||
element: <SettingsPage />,
|
||||
},
|
||||
],
|
||||
},
|
||||
]);
|
||||
|
||||
export function Router() {
|
||||
return <RouterProvider router={router} />;
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { Provider } from 'react-redux';
|
||||
import { initialConfigState } from '@/features/config/slice';
|
||||
import { initialRolloutsUiState } from '@/features/rollouts/slice';
|
||||
import type { AlertsState, AlertTone } from '@/features/ui/alert';
|
||||
import { initialDrawerState } from '@/features/ui/drawer/slice';
|
||||
import { createAppStore } from '@/store';
|
||||
import { STORY_BASE_URL, STORY_DATE_NOW_MS } from '../../.storybook/constants';
|
||||
import { AppAlertBanner } from './AppAlertBanner';
|
||||
|
||||
const meta: Meta<typeof AppAlertBanner> = {
|
||||
title: 'Components/AppAlertBanner',
|
||||
component: AppAlertBanner,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof AppAlertBanner>;
|
||||
|
||||
function renderWithAlert(message: string, tone: AlertTone) {
|
||||
const alertState: AlertsState = {
|
||||
alerts: [
|
||||
{
|
||||
id: 'storybook-alert',
|
||||
message,
|
||||
tone,
|
||||
isVisible: true,
|
||||
createdAt: STORY_DATE_NOW_MS,
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
const store = createAppStore({
|
||||
config: {
|
||||
...initialConfigState,
|
||||
baseUrl: STORY_BASE_URL,
|
||||
},
|
||||
drawer: initialDrawerState,
|
||||
rollouts: initialRolloutsUiState,
|
||||
alert: alertState,
|
||||
});
|
||||
|
||||
return (
|
||||
<Provider store={store}>
|
||||
<div style={{ padding: 24 }}>
|
||||
<AppAlertBanner />
|
||||
</div>
|
||||
</Provider>
|
||||
);
|
||||
}
|
||||
|
||||
export const InfoAlert: Story = {
|
||||
render: () => renderWithAlert('Background synchronization completed successfully.', 'info'),
|
||||
};
|
||||
|
||||
export const WarningAlert: Story = {
|
||||
render: () =>
|
||||
renderWithAlert('Rollout data may be stale. Check your network connection before continuing.', 'warning'),
|
||||
};
|
||||
|
||||
export const ErrorAlert: Story = {
|
||||
render: () =>
|
||||
renderWithAlert('Unable to reach the Agent-lightning API. Retry or adjust the backend settings.', 'error'),
|
||||
};
|
||||
|
||||
export const NoAlert: Story = {
|
||||
render: () => {
|
||||
const store = createAppStore({
|
||||
config: {
|
||||
...initialConfigState,
|
||||
baseUrl: STORY_BASE_URL,
|
||||
},
|
||||
drawer: initialDrawerState,
|
||||
rollouts: initialRolloutsUiState,
|
||||
alert: { alerts: [] },
|
||||
});
|
||||
|
||||
return (
|
||||
<Provider store={store}>
|
||||
<div style={{ padding: 24 }}>
|
||||
<AppAlertBanner />
|
||||
</div>
|
||||
</Provider>
|
||||
);
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,84 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useEffect, useState } from 'react';
|
||||
import { IconAlertCircle, IconAlertTriangle, IconInfoCircle } from '@tabler/icons-react';
|
||||
import { Notification, Portal, Transition } from '@mantine/core';
|
||||
import { hideAlert, selectHighestPriorityAlert, type AppAlert } from '@/features/ui/alert';
|
||||
import { useAppDispatch, useAppSelector } from '@/store/hooks';
|
||||
|
||||
const ALERT_META = {
|
||||
info: {
|
||||
color: 'blue',
|
||||
icon: IconInfoCircle,
|
||||
},
|
||||
warning: {
|
||||
color: 'yellow',
|
||||
icon: IconAlertTriangle,
|
||||
},
|
||||
error: {
|
||||
color: 'red',
|
||||
icon: IconAlertCircle,
|
||||
},
|
||||
} as const;
|
||||
|
||||
export function AppAlertBanner() {
|
||||
const dispatch = useAppDispatch();
|
||||
const alert = useAppSelector(selectHighestPriorityAlert);
|
||||
const [transitionAlert, setTransitionAlert] = useState<AppAlert | null>(alert);
|
||||
|
||||
useEffect(() => {
|
||||
if (alert) {
|
||||
setTransitionAlert(alert);
|
||||
}
|
||||
}, [alert]);
|
||||
|
||||
const handleClose = (id?: string) => {
|
||||
if (id) {
|
||||
dispatch(hideAlert({ id }));
|
||||
}
|
||||
};
|
||||
|
||||
const currentAlert = alert ?? transitionAlert;
|
||||
const mounted = Boolean(alert);
|
||||
|
||||
if (!currentAlert) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const meta = ALERT_META[currentAlert.tone];
|
||||
const IconComponent = meta.icon;
|
||||
|
||||
return (
|
||||
<Portal>
|
||||
<Transition
|
||||
mounted={mounted}
|
||||
transition='slide-down'
|
||||
duration={200}
|
||||
timingFunction='ease'
|
||||
onExited={() => setTransitionAlert(null)}
|
||||
>
|
||||
{(styles) => (
|
||||
<Notification
|
||||
icon={<IconComponent size={18} />}
|
||||
color={meta.color}
|
||||
variant='light'
|
||||
withCloseButton
|
||||
onClose={() => handleClose(currentAlert.id)}
|
||||
style={{
|
||||
position: 'fixed',
|
||||
top: 16,
|
||||
right: 16,
|
||||
maxWidth: 450,
|
||||
width: 'calc(100% - 32px)',
|
||||
zIndex: 2000,
|
||||
boxShadow: 'var(--mantine-shadow-md)',
|
||||
...styles,
|
||||
}}
|
||||
>
|
||||
{currentAlert.message}
|
||||
</Notification>
|
||||
)}
|
||||
</Transition>
|
||||
</Portal>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,583 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo, useRef, useState, type ReactNode } from 'react';
|
||||
import { Editor } from '@monaco-editor/react';
|
||||
import { IconCheck, IconCopy } from '@tabler/icons-react';
|
||||
import type { DataTableSortStatus } from 'mantine-datatable';
|
||||
import { createSearchParams, Link, useInRouterContext, useLocation } from 'react-router-dom';
|
||||
import {
|
||||
ActionIcon,
|
||||
Anchor,
|
||||
Badge,
|
||||
Box,
|
||||
CopyButton,
|
||||
Drawer,
|
||||
Group,
|
||||
Stack,
|
||||
Text,
|
||||
Tooltip,
|
||||
useMantineColorScheme,
|
||||
} from '@mantine/core';
|
||||
import { useGetSpansQuery } from '@/features/rollouts';
|
||||
import { closeDrawer, openDrawer, selectDrawerContent, selectDrawerIsOpen } from '@/features/ui/drawer';
|
||||
import { useAppDispatch, useAppSelector } from '@/store/hooks';
|
||||
import type { Attempt, AttemptStatus, Rollout, RolloutStatus, Span, Worker } from '@/types';
|
||||
import { formatStatusLabel } from '@/utils/format';
|
||||
import { TracesTable, type TracesTableRecord } from './TracesTable.component';
|
||||
|
||||
const ATTEMPT_STATUS_COLORS: Record<AttemptStatus, string> = {
|
||||
failed: 'red',
|
||||
preparing: 'violet',
|
||||
running: 'blue',
|
||||
succeeded: 'teal',
|
||||
timeout: 'orange',
|
||||
unresponsive: 'orange',
|
||||
};
|
||||
|
||||
const ROLLOUT_STATUS_COLORS: Record<RolloutStatus, string> = {
|
||||
cancelled: 'gray',
|
||||
failed: 'red',
|
||||
preparing: 'violet',
|
||||
queuing: 'blue',
|
||||
requeuing: 'cyan',
|
||||
running: 'blue',
|
||||
succeeded: 'teal',
|
||||
};
|
||||
|
||||
const SPAN_STATUS_COLORS: Record<Span['status']['status_code'], string> = {
|
||||
UNSET: 'gray',
|
||||
OK: 'teal',
|
||||
ERROR: 'red',
|
||||
};
|
||||
|
||||
const WORKER_STATUS_COLORS: Record<Worker['status'], string> = {
|
||||
busy: 'orange',
|
||||
idle: 'teal',
|
||||
unknown: 'gray',
|
||||
};
|
||||
|
||||
const TRACES_SORT_FIELD_MAP: Record<string, string> = {
|
||||
name: 'name',
|
||||
traceId: 'trace_id',
|
||||
spanId: 'span_id',
|
||||
parentId: 'parent_id',
|
||||
statusCode: 'status_code',
|
||||
startTime: 'start_time',
|
||||
duration: 'duration',
|
||||
};
|
||||
|
||||
type SortDirection = 'asc' | 'desc';
|
||||
|
||||
type LocalSortState = {
|
||||
column: string;
|
||||
direction: SortDirection;
|
||||
};
|
||||
|
||||
function resolveTracesSortField(column: string): string {
|
||||
return TRACES_SORT_FIELD_MAP[column] ?? 'start_time';
|
||||
}
|
||||
|
||||
function getStatusBadgeColor(status: RolloutStatus | AttemptStatus, isAttempt: boolean) {
|
||||
if (isAttempt) {
|
||||
return ATTEMPT_STATUS_COLORS[status as AttemptStatus] ?? 'gray';
|
||||
}
|
||||
|
||||
return ROLLOUT_STATUS_COLORS[status as RolloutStatus] ?? 'gray';
|
||||
}
|
||||
|
||||
function formatJson(value: unknown) {
|
||||
try {
|
||||
return JSON.stringify(value, null, 2);
|
||||
} catch {
|
||||
return String(value);
|
||||
}
|
||||
}
|
||||
|
||||
export type AppDrawerProps = {
|
||||
opened: boolean;
|
||||
onClose: () => void;
|
||||
title?: ReactNode;
|
||||
body?: ReactNode;
|
||||
};
|
||||
|
||||
export function AppDrawer({ opened, onClose, title, body }: AppDrawerProps) {
|
||||
return (
|
||||
<Drawer
|
||||
position='right'
|
||||
size='lg'
|
||||
opened={opened}
|
||||
onClose={onClose}
|
||||
overlayProps={{ opacity: 0.5 }}
|
||||
withinPortal
|
||||
styles={{
|
||||
content: {
|
||||
display: 'flex',
|
||||
flexDirection: 'column',
|
||||
maxHeight: '100vh',
|
||||
},
|
||||
body: {
|
||||
flex: 1,
|
||||
display: 'flex',
|
||||
flexDirection: 'column',
|
||||
padding: 'var(--mantine-spacing-md)',
|
||||
minHeight: 0,
|
||||
overflow: 'hidden',
|
||||
},
|
||||
}}
|
||||
title={title}
|
||||
>
|
||||
<Stack gap='md' h='100%' style={{ flex: 1, minHeight: 0 }}>
|
||||
{body}
|
||||
</Stack>
|
||||
</Drawer>
|
||||
);
|
||||
}
|
||||
|
||||
type TraceDrawerTitleProps = {
|
||||
span: Span;
|
||||
};
|
||||
|
||||
export function TraceDrawerTitle({ span }: TraceDrawerTitleProps) {
|
||||
const spanStatusCode = span.status?.status_code ?? null;
|
||||
const spanBadgeColor = spanStatusCode ? (SPAN_STATUS_COLORS[spanStatusCode] ?? 'gray') : undefined;
|
||||
|
||||
return (
|
||||
<Stack gap={3}>
|
||||
<Group gap={6}>
|
||||
<Text fw={600}>{span.name ?? span.spanId}</Text>
|
||||
{spanStatusCode ? (
|
||||
<Badge size='sm' variant='light' color={spanBadgeColor}>
|
||||
{spanStatusCode}
|
||||
</Badge>
|
||||
) : null}
|
||||
</Group>
|
||||
<Group gap={6}>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{span.spanId}
|
||||
</Text>
|
||||
<CopyButton value={span.spanId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy span ID ${span.spanId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
<Group gap='xs'>
|
||||
<Group gap={3}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Rollout
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{span.rolloutId}
|
||||
</Text>
|
||||
</Group>
|
||||
<Group gap={3}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Attempt
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{span.attemptId ?? '—'}
|
||||
</Text>
|
||||
</Group>
|
||||
</Group>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
type RolloutAttemptDrawerTitleProps = {
|
||||
rollout: Rollout;
|
||||
attempt: Attempt | null;
|
||||
};
|
||||
|
||||
export function RolloutAttemptDrawerTitle({ rollout, attempt }: RolloutAttemptDrawerTitleProps) {
|
||||
const rolloutId = rollout.rolloutId;
|
||||
const attemptId = attempt?.attemptId ?? null;
|
||||
const rolloutStatus = rollout.status ?? null;
|
||||
const attemptStatus = attempt?.status ?? null;
|
||||
const rolloutStatusLabel = rolloutStatus ? formatStatusLabel(rolloutStatus) : null;
|
||||
const attemptStatusLabel = attemptStatus ? formatStatusLabel(attemptStatus) : null;
|
||||
const hasStatusMismatch = rolloutStatus !== null && attemptStatus !== null && rolloutStatus !== attemptStatus;
|
||||
const rolloutBadgeColor = rolloutStatus ? getStatusBadgeColor(rolloutStatus, false) : undefined;
|
||||
const attemptBadgeColor = attemptStatus ? getStatusBadgeColor(attemptStatus, true) : undefined;
|
||||
const showRolloutBadgeInHeading = Boolean(rolloutStatusLabel && (!attemptStatus || hasStatusMismatch));
|
||||
const showAttemptBadge = Boolean(attemptStatusLabel && attemptStatus);
|
||||
|
||||
return (
|
||||
<Stack gap={3}>
|
||||
<Group gap={6}>
|
||||
<Text fw={600}>{rolloutId}</Text>
|
||||
<CopyButton value={rolloutId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy rollout ID ${rolloutId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
{showRolloutBadgeInHeading && rolloutStatusLabel ? (
|
||||
<Badge size='sm' variant='light' color={rolloutBadgeColor}>
|
||||
{rolloutStatusLabel}
|
||||
</Badge>
|
||||
) : null}
|
||||
</Group>
|
||||
<Group gap='xs'>
|
||||
{attemptId ? (
|
||||
<Group gap={3}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Attempt
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{attemptId}
|
||||
</Text>
|
||||
</Group>
|
||||
) : null}
|
||||
{showAttemptBadge && attemptStatusLabel ? (
|
||||
<Badge size='sm' variant='light' color={attemptBadgeColor}>
|
||||
{attemptStatusLabel}
|
||||
</Badge>
|
||||
) : null}
|
||||
{!showRolloutBadgeInHeading && !attemptStatus && rolloutStatusLabel ? (
|
||||
<Badge size='sm' variant='light' color={rolloutBadgeColor}>
|
||||
{rolloutStatusLabel}
|
||||
</Badge>
|
||||
) : null}
|
||||
</Group>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
type JsonEditorProps = {
|
||||
value: unknown;
|
||||
};
|
||||
|
||||
export function JsonEditor({ value }: JsonEditorProps) {
|
||||
const { colorScheme } = useMantineColorScheme();
|
||||
const editorTheme = colorScheme === 'dark' ? 'vs-dark' : 'vs-light';
|
||||
|
||||
return (
|
||||
<Box data-testid='json-editor-container' style={{ flex: 1, minHeight: 0 }}>
|
||||
<Editor
|
||||
height='100%'
|
||||
language='json'
|
||||
value={formatJson(value)}
|
||||
theme={editorTheme}
|
||||
options={{
|
||||
readOnly: true,
|
||||
domReadOnly: true,
|
||||
minimap: { enabled: false },
|
||||
automaticLayout: true,
|
||||
scrollBeyondLastLine: false,
|
||||
fontSize: 13,
|
||||
}}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
type RolloutTracesDrawerBodyProps = {
|
||||
rollout: Rollout;
|
||||
attempt: Attempt | null;
|
||||
onShowRollout: (record: TracesTableRecord) => void;
|
||||
onShowSpanDetail: (record: TracesTableRecord) => void;
|
||||
};
|
||||
|
||||
function RolloutTracesDrawerBody({ rollout, attempt, onShowRollout, onShowSpanDetail }: RolloutTracesDrawerBodyProps) {
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(100);
|
||||
const [sort, setSort] = useState<LocalSortState>({
|
||||
column: 'startTime',
|
||||
direction: 'desc',
|
||||
});
|
||||
|
||||
useEffect(() => {
|
||||
setPage(1);
|
||||
}, [rollout.rolloutId, attempt?.attemptId]);
|
||||
|
||||
const queryArgs = useMemo(
|
||||
() => ({
|
||||
rolloutId: rollout.rolloutId,
|
||||
attemptId: attempt?.attemptId ?? undefined,
|
||||
limit: recordsPerPage,
|
||||
offset: Math.max(0, (page - 1) * recordsPerPage),
|
||||
sortBy: resolveTracesSortField(sort.column),
|
||||
sortOrder: sort.direction,
|
||||
}),
|
||||
[rollout.rolloutId, attempt?.attemptId, recordsPerPage, page, sort],
|
||||
);
|
||||
|
||||
const { data, isFetching, isError, error, refetch } = useGetSpansQuery(queryArgs);
|
||||
const spans = data?.items ?? [];
|
||||
const totalRecords = data?.total ?? 0;
|
||||
const tracesLinkSearch = useMemo(() => {
|
||||
const params = createSearchParams({
|
||||
rolloutId: rollout.rolloutId,
|
||||
...(attempt?.attemptId ? { attemptId: attempt.attemptId } : {}),
|
||||
});
|
||||
return params.toString();
|
||||
}, [attempt?.attemptId, rollout.rolloutId]);
|
||||
const tracesLinkHref = tracesLinkSearch ? `/traces?${tracesLinkSearch}` : '/traces';
|
||||
const isWithinRouter = useInRouterContext();
|
||||
|
||||
const handleSortStatusChange = useCallback((status: DataTableSortStatus<TracesTableRecord>) => {
|
||||
setSort({
|
||||
column: status.columnAccessor as string,
|
||||
direction: status.direction,
|
||||
});
|
||||
}, []);
|
||||
|
||||
const handlePageChange = useCallback((nextPage: number) => {
|
||||
setPage(nextPage);
|
||||
}, []);
|
||||
|
||||
const handleRecordsPerPageChange = useCallback((value: number) => {
|
||||
setRecordsPerPage(value);
|
||||
setPage(1);
|
||||
}, []);
|
||||
|
||||
return (
|
||||
<Stack gap='md' style={{ flex: 1, minHeight: 0 }}>
|
||||
<Group justify='space-between' align='center' gap='sm' wrap='nowrap'>
|
||||
<Text size='sm' style={{ flex: 1, minWidth: 0 }}>
|
||||
Showing spans for{' '}
|
||||
<Text component='span' fw={600}>
|
||||
{rollout.rolloutId}
|
||||
{attempt ? ` · Attempt ${attempt.sequenceId} (${attempt.attemptId})` : ' · Latest attempt'}
|
||||
</Text>
|
||||
</Text>
|
||||
{isWithinRouter ? (
|
||||
<Anchor
|
||||
component={Link}
|
||||
to={tracesLinkHref}
|
||||
size='sm'
|
||||
aria-label={`Open traces page for rollout ${rollout.rolloutId}${
|
||||
attempt ? ` attempt ${attempt.sequenceId}` : ''
|
||||
}`}
|
||||
>
|
||||
View full traces
|
||||
</Anchor>
|
||||
) : (
|
||||
<Anchor
|
||||
href={tracesLinkHref}
|
||||
size='sm'
|
||||
aria-label={`Open traces page for rollout ${rollout.rolloutId}${
|
||||
attempt ? ` attempt ${attempt.sequenceId}` : ''
|
||||
}`}
|
||||
>
|
||||
View full traces
|
||||
</Anchor>
|
||||
)}
|
||||
</Group>
|
||||
<Box data-testid='traces-drawer-table-container' style={{ flex: 1, minHeight: 0, overflow: 'auto' }}>
|
||||
<TracesTable
|
||||
spans={spans}
|
||||
totalRecords={totalRecords}
|
||||
isFetching={isFetching}
|
||||
isError={isError}
|
||||
error={error}
|
||||
searchTerm=''
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
onPageChange={handlePageChange}
|
||||
onRecordsPerPageChange={handleRecordsPerPageChange}
|
||||
onResetFilters={() => {}}
|
||||
onRefetch={refetch}
|
||||
onShowRollout={onShowRollout}
|
||||
onShowSpanDetail={onShowSpanDetail}
|
||||
recordsPerPageOptions={[50, 100, 200, 500]}
|
||||
/>
|
||||
</Box>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
type WorkerDrawerTitleProps = {
|
||||
worker: Worker;
|
||||
};
|
||||
|
||||
function WorkerDrawerTitle({ worker }: WorkerDrawerTitleProps) {
|
||||
const badgeColor = WORKER_STATUS_COLORS[worker.status] ?? 'gray';
|
||||
return (
|
||||
<Stack gap={3}>
|
||||
<Group gap={6} align='center'>
|
||||
<Text fw={600}>{worker.workerId}</Text>
|
||||
<CopyButton value={worker.workerId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy worker ID ${worker.workerId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
<Badge size='sm' variant='light' color={badgeColor}>
|
||||
{formatStatusLabel(worker.status)}
|
||||
</Badge>
|
||||
</Group>
|
||||
<Group gap='xl'>
|
||||
<Group gap={4}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Rollout
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{worker.currentRolloutId ?? '—'}
|
||||
</Text>
|
||||
</Group>
|
||||
<Group gap={4}>
|
||||
<Text size='sm' c='dimmed' fw={500}>
|
||||
Attempt
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed'>
|
||||
{worker.currentAttemptId ?? '—'}
|
||||
</Text>
|
||||
</Group>
|
||||
</Group>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
export function AppDrawerContainer() {
|
||||
const dispatch = useAppDispatch();
|
||||
const isOpen = useAppSelector(selectDrawerIsOpen);
|
||||
const content = useAppSelector(selectDrawerContent);
|
||||
const isRouterAvailable = useInRouterContext();
|
||||
|
||||
const handleClose = useCallback(() => {
|
||||
dispatch(closeDrawer());
|
||||
}, [dispatch]);
|
||||
const handleNavigation = useCallback(() => {
|
||||
if (isOpen) {
|
||||
dispatch(closeDrawer());
|
||||
}
|
||||
}, [dispatch, isOpen]);
|
||||
|
||||
const derivedContent = useMemo(() => {
|
||||
if (!content) {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (content.type === 'worker-detail') {
|
||||
const { worker } = content;
|
||||
const title = <WorkerDrawerTitle worker={worker} />;
|
||||
const body = <JsonEditor value={worker} />;
|
||||
return { title, body };
|
||||
}
|
||||
|
||||
if (content.type === 'trace-detail') {
|
||||
const { span } = content;
|
||||
const title = <TraceDrawerTitle span={span} />;
|
||||
const body = <JsonEditor value={span} />;
|
||||
|
||||
return { title, body };
|
||||
}
|
||||
|
||||
const rollout = content.rollout;
|
||||
const attempt = content.attempt;
|
||||
const title = <RolloutAttemptDrawerTitle rollout={rollout} attempt={attempt} />;
|
||||
|
||||
if (content.type === 'rollout-json') {
|
||||
const jsonValue = content.isNested && content.attempt ? content.attempt : rollout;
|
||||
const body = jsonValue ? <JsonEditor value={jsonValue} /> : null;
|
||||
return { title, body };
|
||||
}
|
||||
|
||||
if (content.type === 'rollout-traces') {
|
||||
const body = (
|
||||
<RolloutTracesDrawerBody
|
||||
rollout={rollout}
|
||||
attempt={attempt}
|
||||
onShowRollout={() => {
|
||||
const attemptForRecord = attempt ?? rollout.attempt ?? null;
|
||||
dispatch(
|
||||
openDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout,
|
||||
attempt: attemptForRecord,
|
||||
isNested: content.isNested,
|
||||
}),
|
||||
);
|
||||
}}
|
||||
onShowSpanDetail={(record) => {
|
||||
const attemptForRecord = attempt ?? rollout.attempt ?? null;
|
||||
dispatch(
|
||||
openDrawer({
|
||||
type: 'trace-detail',
|
||||
span: record,
|
||||
rollout,
|
||||
attempt: attemptForRecord,
|
||||
}),
|
||||
);
|
||||
}}
|
||||
/>
|
||||
);
|
||||
|
||||
return { title, body };
|
||||
}
|
||||
|
||||
return null;
|
||||
}, [content, dispatch]);
|
||||
|
||||
if (!content || !derivedContent) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const { title, body } = derivedContent;
|
||||
|
||||
return (
|
||||
<>
|
||||
{isRouterAvailable ? <DrawerLocationWatcher onNavigation={handleNavigation} /> : null}
|
||||
<AppDrawer opened={isOpen} onClose={handleClose} title={title} body={body} />
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
||||
type DrawerLocationWatcherProps = {
|
||||
onNavigation: () => void;
|
||||
};
|
||||
|
||||
function DrawerLocationWatcher({ onNavigation }: DrawerLocationWatcherProps) {
|
||||
const location = useLocation();
|
||||
const lastLocationKeyRef = useRef(location.key);
|
||||
|
||||
useEffect(() => {
|
||||
if (lastLocationKeyRef.current === location.key) {
|
||||
return;
|
||||
}
|
||||
lastLocationKeyRef.current = location.key;
|
||||
onNavigation();
|
||||
}, [location.key, onNavigation]);
|
||||
|
||||
return null;
|
||||
}
|
||||
@@ -0,0 +1,258 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { Provider } from 'react-redux';
|
||||
import { initialConfigState } from '@/features/config/slice';
|
||||
import { initialResourcesUiState } from '@/features/resources/slice';
|
||||
import { rolloutsApi } from '@/features/rollouts';
|
||||
import { initialRolloutsUiState } from '@/features/rollouts/slice';
|
||||
import { initialTracesUiState } from '@/features/traces/slice';
|
||||
import type { DrawerContent } from '@/features/ui/drawer';
|
||||
import { createAppStore } from '@/store';
|
||||
import type { Attempt, Rollout, Span } from '@/types';
|
||||
import { STORY_BASE_URL, STORY_DATE_NOW_SECONDS } from '../../.storybook/constants';
|
||||
import { AppDrawerContainer } from './AppDrawer.component';
|
||||
|
||||
const meta = {
|
||||
title: 'Components/AppDrawer',
|
||||
component: AppDrawerContainer,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
} satisfies Meta<typeof AppDrawerContainer>;
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof AppDrawerContainer>;
|
||||
|
||||
const now = STORY_DATE_NOW_SECONDS;
|
||||
|
||||
const baseAttempt: Attempt = {
|
||||
rolloutId: 'ro-story-001',
|
||||
attemptId: 'at-story-001',
|
||||
sequenceId: 1,
|
||||
startTime: now - 3600,
|
||||
endTime: null,
|
||||
status: 'running',
|
||||
workerId: 'worker-story',
|
||||
lastHeartbeatTime: now - 42,
|
||||
metadata: { info: 'Sample metadata', runId: 'run-123' },
|
||||
};
|
||||
|
||||
const baseRollout: Rollout = {
|
||||
rolloutId: 'ro-story-001',
|
||||
input: {
|
||||
task: 'Generate daily summary',
|
||||
payload: { account: 'enterprise', date: '2024-02-19' },
|
||||
},
|
||||
startTime: now - 4000,
|
||||
endTime: null,
|
||||
mode: 'train',
|
||||
resourcesId: 'rs-story-001',
|
||||
status: 'running',
|
||||
config: { retries: 1, priority: 'high' },
|
||||
metadata: { owner: 'storybook' },
|
||||
attempt: baseAttempt,
|
||||
};
|
||||
|
||||
const noAttemptRollout: Rollout = {
|
||||
...baseRollout,
|
||||
status: 'queuing',
|
||||
attempt: null,
|
||||
};
|
||||
|
||||
const mismatchRollout: Rollout = {
|
||||
...baseRollout,
|
||||
status: 'running',
|
||||
attempt: {
|
||||
...baseAttempt,
|
||||
status: 'failed',
|
||||
endTime: now - 1200,
|
||||
metadata: { info: 'Latest attempt failed', reason: 'Timeout' },
|
||||
},
|
||||
};
|
||||
|
||||
const sampleSpan: Span = {
|
||||
rolloutId: 'ro-story-001',
|
||||
attemptId: 'at-story-001',
|
||||
sequenceId: 2,
|
||||
traceId: 'tr-story-001',
|
||||
spanId: 'sp-story-001',
|
||||
parentId: null,
|
||||
name: 'Fetch Resources',
|
||||
status: { status_code: 'OK', description: 'Completed successfully' },
|
||||
attributes: {
|
||||
'http.method': 'GET',
|
||||
'http.url': 'https://api.example.com/resources',
|
||||
'duration_ms': 120,
|
||||
},
|
||||
startTime: now - 240,
|
||||
endTime: now - 120,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
};
|
||||
|
||||
const sampleTraces: Span[] = [
|
||||
sampleSpan,
|
||||
{
|
||||
...sampleSpan,
|
||||
spanId: 'sp-story-002',
|
||||
name: 'Process Response',
|
||||
parentId: 'sp-story-001',
|
||||
sequenceId: 3,
|
||||
status: { status_code: 'ERROR', description: 'Unexpected response code' },
|
||||
attributes: {
|
||||
...sampleSpan.attributes,
|
||||
duration_ms: 240,
|
||||
},
|
||||
startTime: now - 120,
|
||||
endTime: now - 30,
|
||||
},
|
||||
];
|
||||
|
||||
function renderWithDrawer(content: DrawerContent, options?: { spans?: Span[] }) {
|
||||
const store = createAppStore({
|
||||
config: {
|
||||
...initialConfigState,
|
||||
baseUrl: STORY_BASE_URL,
|
||||
},
|
||||
rollouts: initialRolloutsUiState,
|
||||
resources: initialResourcesUiState,
|
||||
traces: initialTracesUiState,
|
||||
drawer: {
|
||||
isOpen: true,
|
||||
content,
|
||||
},
|
||||
});
|
||||
|
||||
if (content.type === 'rollout-traces' && options?.spans) {
|
||||
const defaultLimit = 100;
|
||||
const queryArgs = {
|
||||
rolloutId: content.rollout.rolloutId,
|
||||
attemptId: content.attempt?.attemptId ?? undefined,
|
||||
limit: defaultLimit,
|
||||
offset: 0,
|
||||
sortBy: 'start_time',
|
||||
sortOrder: 'desc' as const,
|
||||
};
|
||||
|
||||
store.dispatch(
|
||||
rolloutsApi.util.upsertQueryData('getSpans', queryArgs, {
|
||||
items: options.spans,
|
||||
total: options.spans.length,
|
||||
limit: defaultLimit,
|
||||
offset: 0,
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<Provider store={store}>
|
||||
<AppDrawerContainer />
|
||||
</Provider>
|
||||
);
|
||||
}
|
||||
|
||||
export const RolloutJson: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout: baseRollout,
|
||||
attempt: baseRollout.attempt,
|
||||
isNested: false,
|
||||
}),
|
||||
};
|
||||
|
||||
export const NestedAttemptJson: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout: baseRollout,
|
||||
attempt: {
|
||||
...baseAttempt,
|
||||
attemptId: 'at-story-002',
|
||||
sequenceId: 2,
|
||||
status: 'failed',
|
||||
endTime: now - 1200,
|
||||
metadata: { info: 'Secondary attempt', reason: 'Timeout' },
|
||||
},
|
||||
isNested: true,
|
||||
}),
|
||||
};
|
||||
|
||||
export const RolloutTraces: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer(
|
||||
{
|
||||
type: 'rollout-traces',
|
||||
rollout: baseRollout,
|
||||
attempt: baseRollout.attempt,
|
||||
isNested: false,
|
||||
},
|
||||
{ spans: sampleTraces },
|
||||
),
|
||||
};
|
||||
|
||||
export const NoAttempt: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout: noAttemptRollout,
|
||||
attempt: null,
|
||||
isNested: false,
|
||||
}),
|
||||
};
|
||||
|
||||
export const StatusMismatch: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout: mismatchRollout,
|
||||
attempt: mismatchRollout.attempt,
|
||||
isNested: false,
|
||||
}),
|
||||
};
|
||||
|
||||
export const SpanDetail: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'trace-detail',
|
||||
span: sampleSpan,
|
||||
rollout: mismatchRollout,
|
||||
attempt: mismatchRollout.attempt,
|
||||
}),
|
||||
};
|
||||
|
||||
export const LightTheme: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'rollout-json',
|
||||
rollout: baseRollout,
|
||||
attempt: baseRollout.attempt,
|
||||
isNested: false,
|
||||
}),
|
||||
parameters: {
|
||||
theme: 'light',
|
||||
},
|
||||
};
|
||||
|
||||
export const DarkTheme: Story = {
|
||||
render: () =>
|
||||
renderWithDrawer({
|
||||
type: 'trace-detail',
|
||||
span: {
|
||||
...sampleSpan,
|
||||
spanId: 'sp-story-002',
|
||||
name: 'Process Response',
|
||||
status: { status_code: 'ERROR', description: 'Unexpected response code' },
|
||||
},
|
||||
rollout: mismatchRollout,
|
||||
attempt: mismatchRollout.attempt,
|
||||
}),
|
||||
parameters: {
|
||||
theme: 'dark',
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,322 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo, useState, type ReactNode, type SetStateAction } from 'react';
|
||||
import { IconCheck, IconCopy, IconRefresh } from '@tabler/icons-react';
|
||||
import { DataTable, type DataTableColumn, type DataTableSortStatus } from 'mantine-datatable';
|
||||
import { ActionIcon, Box, Button, CopyButton, Group, Stack, Text, Tooltip } from '@mantine/core';
|
||||
import { useElementSize, useViewportSize } from '@mantine/hooks';
|
||||
import { getLayoutAwareWidth } from '@/layouts/helper';
|
||||
import type { Resources } from '@/types';
|
||||
import { getErrorDescriptor } from '@/utils/error';
|
||||
import { formatDateTime, safeStringify } from '@/utils/format';
|
||||
import { createResponsiveColumns, type ColumnVisibilityConfig } from '@/utils/table';
|
||||
|
||||
const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
resourcesId: { fixedWidth: 12, priority: 0 },
|
||||
version: { fixedWidth: 8, priority: 1 },
|
||||
createTime: { fixedWidth: 14, priority: 2 },
|
||||
updateTime: { fixedWidth: 14, priority: 2 },
|
||||
resourceCount: { fixedWidth: 8, priority: 3 },
|
||||
resourcesPreview: { minWidth: 16, priority: 4 },
|
||||
};
|
||||
|
||||
export type ResourcesTableRecord = Resources & {
|
||||
resourceCount: number;
|
||||
canExpand: boolean;
|
||||
resourcesPreview: string;
|
||||
};
|
||||
|
||||
function buildResourcesRecord(resources: Resources): ResourcesTableRecord {
|
||||
const resourceCount = Object.keys(resources.resources ?? {}).length;
|
||||
const resourcesValue =
|
||||
resources.resources === null || typeof resources.resources === 'undefined'
|
||||
? '—'
|
||||
: typeof resources.resources === 'string'
|
||||
? resources.resources
|
||||
: safeStringify(resources.resources);
|
||||
|
||||
return {
|
||||
...resources,
|
||||
resourceCount,
|
||||
canExpand: resourceCount > 0,
|
||||
resourcesPreview: resourcesValue,
|
||||
};
|
||||
}
|
||||
|
||||
type ResourcesColumnsOptions = Record<string, never>;
|
||||
|
||||
function createResourcesColumns(_options: ResourcesColumnsOptions): DataTableColumn<ResourcesTableRecord>[] {
|
||||
return [
|
||||
{
|
||||
accessor: 'resourcesId',
|
||||
title: 'Resources ID',
|
||||
sortable: true,
|
||||
render: ({ resourcesId }) => (
|
||||
<Group gap={2}>
|
||||
<Text fw={500} size='sm'>
|
||||
{resourcesId}
|
||||
</Text>
|
||||
<CopyButton value={resourcesId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy resources ID ${resourcesId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'version',
|
||||
title: 'Version',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ version }) => <Text size='sm'>{version}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'createTime',
|
||||
title: 'Created',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ createTime }) => <Text size='sm'>{formatDateTime(createTime)}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'updateTime',
|
||||
title: 'Updated',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ updateTime }) => <Text size='sm'>{formatDateTime(updateTime)}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'resourceCount',
|
||||
title: 'Count',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ resourceCount }) => <Text size='sm'>{resourceCount}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'resourcesPreview',
|
||||
title: 'Preview',
|
||||
render: ({ resourcesPreview }) => (
|
||||
<Text size='sm' ff='monospace' c='dimmed' lineClamp={1} style={{ width: '100%' }}>
|
||||
{resourcesPreview}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
type RowExpansionRenderer = (context: {
|
||||
resources: Resources;
|
||||
columns: DataTableColumn<ResourcesTableRecord>[];
|
||||
}) => ReactNode;
|
||||
|
||||
export type ResourcesTableProps = {
|
||||
resourcesList: Resources[] | undefined;
|
||||
totalRecords: number;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
error: unknown;
|
||||
searchTerm: string;
|
||||
sort: { column: string; direction: 'asc' | 'desc' };
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
onSortStatusChange: (status: DataTableSortStatus<ResourcesTableRecord>) => void;
|
||||
onPageChange: (page: number) => void;
|
||||
onRecordsPerPageChange: (value: number) => void;
|
||||
onResetFilters: () => void;
|
||||
onRefetch: () => void;
|
||||
recordsPerPageOptions?: number[];
|
||||
renderRowExpansion?: RowExpansionRenderer;
|
||||
};
|
||||
|
||||
export function ResourcesTable({
|
||||
resourcesList,
|
||||
totalRecords,
|
||||
isFetching,
|
||||
isError,
|
||||
error,
|
||||
searchTerm,
|
||||
sort,
|
||||
page,
|
||||
recordsPerPage,
|
||||
onSortStatusChange,
|
||||
onPageChange,
|
||||
onRecordsPerPageChange,
|
||||
onResetFilters,
|
||||
onRefetch,
|
||||
recordsPerPageOptions = DEFAULT_RECORDS_PER_PAGE_OPTIONS,
|
||||
renderRowExpansion,
|
||||
}: ResourcesTableProps) {
|
||||
const [expandedRecordIds, setExpandedRecordIds] = useState<string[]>([]);
|
||||
const { ref: tableContainerRef, width: containerWidth } = useElementSize();
|
||||
const { width: viewportWidth } = useViewportSize();
|
||||
|
||||
const layoutAwareContainerWidth = useMemo(
|
||||
() => getLayoutAwareWidth(containerWidth, viewportWidth),
|
||||
[containerWidth, viewportWidth],
|
||||
);
|
||||
|
||||
const resourcesRecords = useMemo<ResourcesTableRecord[]>(() => {
|
||||
if (!resourcesList) {
|
||||
return [];
|
||||
}
|
||||
return resourcesList.map((resourcesItem) => buildResourcesRecord(resourcesItem));
|
||||
}, [resourcesList]);
|
||||
|
||||
const columns = useMemo(() => createResourcesColumns({}), []);
|
||||
|
||||
const responsiveColumns = useMemo(
|
||||
() => createResponsiveColumns(columns, layoutAwareContainerWidth, COLUMN_VISIBILITY),
|
||||
[columns, layoutAwareContainerWidth],
|
||||
);
|
||||
|
||||
const totalPages = useMemo(
|
||||
() => Math.max(1, Math.ceil(Math.max(0, totalRecords) / Math.max(1, recordsPerPage))),
|
||||
[recordsPerPage, totalRecords],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (page > totalPages) {
|
||||
onPageChange(totalPages);
|
||||
}
|
||||
}, [onPageChange, page, totalPages]);
|
||||
|
||||
useEffect(() => {
|
||||
setExpandedRecordIds((current) =>
|
||||
current.filter((id) => resourcesRecords.some((record) => record.resourcesId === id && record.canExpand)),
|
||||
);
|
||||
}, [resourcesRecords]);
|
||||
|
||||
const hasActiveFilters = searchTerm.trim().length > 0;
|
||||
|
||||
const sortStatus: DataTableSortStatus<ResourcesTableRecord> = {
|
||||
columnAccessor: sort.column,
|
||||
direction: sort.direction,
|
||||
};
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<ResourcesTableRecord>) => {
|
||||
onSortStatusChange(status);
|
||||
},
|
||||
[onSortStatusChange],
|
||||
);
|
||||
|
||||
const errorDescriptor = isError ? getErrorDescriptor(error) : null;
|
||||
const errorMessage = isError
|
||||
? `Resources are temporarily unavailable${errorDescriptor ? ` (${errorDescriptor})` : ''}.`
|
||||
: 'Resources are temporarily unavailable.';
|
||||
|
||||
const emptyState = (
|
||||
<Stack gap='sm' align='center' py='lg'>
|
||||
{isError ? (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
{errorMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Use the retry button to try again, or adjust the filters to broaden the results.
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' color='gray' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Retry
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
No resources found
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
{hasActiveFilters
|
||||
? 'Try adjusting the search to see more results.'
|
||||
: 'Try refreshing to fetch the latest resources.'}
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Refresh
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
|
||||
return (
|
||||
<Box ref={tableContainerRef}>
|
||||
<DataTable<ResourcesTableRecord>
|
||||
classNames={{ root: 'resources-table' }}
|
||||
withTableBorder
|
||||
withColumnBorders
|
||||
highlightOnHover
|
||||
verticalAlign='center'
|
||||
minHeight={resourcesRecords.length === 0 ? 500 : undefined}
|
||||
idAccessor='resourcesId'
|
||||
records={resourcesRecords}
|
||||
columns={responsiveColumns}
|
||||
totalRecords={totalRecords}
|
||||
recordsPerPage={recordsPerPage}
|
||||
page={page}
|
||||
onPageChange={onPageChange}
|
||||
onRecordsPerPageChange={onRecordsPerPageChange}
|
||||
recordsPerPageOptions={recordsPerPageOptions}
|
||||
sortStatus={sortStatus}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
emptyState={resourcesRecords.length === 0 ? emptyState : undefined}
|
||||
rowExpansion={
|
||||
renderRowExpansion
|
||||
? {
|
||||
allowMultiple: true,
|
||||
expandable: ({ record }) => record.canExpand,
|
||||
expanded: {
|
||||
recordIds: expandedRecordIds,
|
||||
onRecordIdsChange: (nextRecordIds: SetStateAction<string[]>) => {
|
||||
setExpandedRecordIds((previous) => {
|
||||
const resolved =
|
||||
typeof nextRecordIds === 'function'
|
||||
? nextRecordIds(previous)
|
||||
: ((nextRecordIds ?? []) as (string | number)[]);
|
||||
return resolved
|
||||
.map(String)
|
||||
.filter((id) =>
|
||||
resourcesRecords.some(
|
||||
(tableRecord) => tableRecord.resourcesId === id && tableRecord.canExpand,
|
||||
),
|
||||
);
|
||||
});
|
||||
},
|
||||
},
|
||||
content: ({ record }) => renderRowExpansion({ resources: record, columns: responsiveColumns }),
|
||||
}
|
||||
: undefined
|
||||
}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo, useState } from 'react';
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import { Box, Stack, TextInput, Title } from '@mantine/core';
|
||||
import type { Resources } from '@/types';
|
||||
import { ResourcesTable } from './ResourcesTable.component';
|
||||
import { ResourcesTree } from './ResourcesTree.component';
|
||||
|
||||
const meta: Meta<typeof ResourcesTable> = {
|
||||
title: 'Components/ResourcesTable',
|
||||
component: ResourcesTable,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof ResourcesTable>;
|
||||
|
||||
const sampleResources: Resources[] = [
|
||||
{
|
||||
resourcesId: 'rs-story-001',
|
||||
version: 1,
|
||||
createTime: 1710806400,
|
||||
updateTime: 1713412800,
|
||||
resources: {
|
||||
model: {
|
||||
name: 'gpt-4',
|
||||
version: '2024-01-01',
|
||||
temperature: 0.7,
|
||||
maxTokens: 2048,
|
||||
topP: 0.9,
|
||||
},
|
||||
database: {
|
||||
host: 'db.example.com',
|
||||
port: 5432,
|
||||
name: 'production',
|
||||
pool: {
|
||||
min: 2,
|
||||
max: 10,
|
||||
idle: 30000,
|
||||
},
|
||||
},
|
||||
cache: {
|
||||
type: 'redis',
|
||||
host: 'cache.example.com',
|
||||
port: 6379,
|
||||
ttl: 3600,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
resourcesId: 'rs-story-002',
|
||||
version: 2,
|
||||
createTime: 1712217600,
|
||||
updateTime: 1714823200,
|
||||
resources: {
|
||||
model: {
|
||||
name: 'claude-3-opus',
|
||||
version: '2024-02-01',
|
||||
temperature: 0.5,
|
||||
maxTokens: 4096,
|
||||
},
|
||||
storage: {
|
||||
type: 's3',
|
||||
bucket: 'training-data',
|
||||
region: 'us-east-1',
|
||||
credentials: {
|
||||
accessKeyId: 'AKIA***',
|
||||
encrypted: true,
|
||||
},
|
||||
},
|
||||
compute: {
|
||||
instances: [
|
||||
{ id: 'i-001', type: 't3.large', zone: 'us-east-1a' },
|
||||
{ id: 'i-002', type: 't3.large', zone: 'us-east-1b' },
|
||||
],
|
||||
autoScaling: {
|
||||
min: 2,
|
||||
max: 10,
|
||||
targetCpu: 70,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
resourcesId: 'rs-story-003',
|
||||
version: 3,
|
||||
createTime: 1709251200,
|
||||
updateTime: 1711856800,
|
||||
resources: {
|
||||
model: {
|
||||
name: 'gpt-3.5-turbo',
|
||||
version: '2023-12-01',
|
||||
temperature: 0.8,
|
||||
maxTokens: 1024,
|
||||
},
|
||||
monitoring: {
|
||||
enabled: true,
|
||||
interval: 60,
|
||||
metrics: ['cpu', 'memory', 'disk', 'network'],
|
||||
alerts: {
|
||||
email: 'ops@example.com',
|
||||
slack: '#alerts',
|
||||
pagerduty: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
resourcesId: 'rs-story-004',
|
||||
version: 1,
|
||||
createTime: 1706745600,
|
||||
updateTime: 1709347200,
|
||||
resources: {
|
||||
apiKeys: {
|
||||
openai: 'sk-***',
|
||||
anthropic: 'sk-ant-***',
|
||||
replicate: 'r8-***',
|
||||
},
|
||||
rateLimits: {
|
||||
requestsPerMinute: 100,
|
||||
tokensPerDay: 1000000,
|
||||
concurrent: 5,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
resourcesId: 'rs-story-005',
|
||||
version: 1,
|
||||
createTime: 1704067200,
|
||||
updateTime: 1706668800,
|
||||
resources: {},
|
||||
},
|
||||
];
|
||||
|
||||
type WrapperProps = {
|
||||
maxWidth: number;
|
||||
resourcesList?: Resources[] | undefined;
|
||||
isFetching?: boolean;
|
||||
isError?: boolean;
|
||||
error?: unknown;
|
||||
};
|
||||
|
||||
function ResourcesTableStoryWrapper({
|
||||
maxWidth,
|
||||
resourcesList = sampleResources,
|
||||
isFetching = false,
|
||||
isError = false,
|
||||
error = null,
|
||||
}: WrapperProps) {
|
||||
const [searchTerm, setSearchTerm] = useState('');
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(5);
|
||||
const [sort, setSort] = useState<{ column: string; direction: 'asc' | 'desc' }>({
|
||||
column: 'resourcesId',
|
||||
direction: 'asc',
|
||||
});
|
||||
|
||||
const baseResources = resourcesList ?? [];
|
||||
|
||||
const filteredResources = useMemo(() => {
|
||||
const normalized = searchTerm.trim().toLowerCase();
|
||||
if (normalized.length === 0) {
|
||||
return baseResources;
|
||||
}
|
||||
return baseResources.filter((resource) => resource.resourcesId.toLowerCase().includes(normalized));
|
||||
}, [baseResources, searchTerm]);
|
||||
|
||||
const sortedResources = useMemo(() => {
|
||||
const items = filteredResources.slice();
|
||||
const resolveSortValue = (resource: Resources, column: string) => {
|
||||
switch (column) {
|
||||
case 'version':
|
||||
return resource.version;
|
||||
case 'createTime':
|
||||
return resource.createTime;
|
||||
case 'updateTime':
|
||||
return resource.updateTime;
|
||||
case 'resourceCount':
|
||||
return Object.keys(resource.resources ?? {}).length;
|
||||
case 'resourcesId':
|
||||
default:
|
||||
return resource.resourcesId;
|
||||
}
|
||||
};
|
||||
|
||||
items.sort((a, b) => {
|
||||
const aValue = resolveSortValue(a, sort.column);
|
||||
const bValue = resolveSortValue(b, sort.column);
|
||||
|
||||
if (aValue === bValue) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
if (typeof aValue === 'number' && typeof bValue === 'number') {
|
||||
return aValue - bValue;
|
||||
}
|
||||
|
||||
return String(aValue).localeCompare(String(bValue));
|
||||
});
|
||||
|
||||
if (sort.direction === 'desc') {
|
||||
items.reverse();
|
||||
}
|
||||
|
||||
return items;
|
||||
}, [filteredResources, sort]);
|
||||
|
||||
const totalRecordsValue = sortedResources.length;
|
||||
|
||||
const pagedResources = useMemo(() => {
|
||||
const startIndex = (page - 1) * recordsPerPage;
|
||||
return sortedResources.slice(startIndex, startIndex + recordsPerPage);
|
||||
}, [page, recordsPerPage, sortedResources]);
|
||||
|
||||
return (
|
||||
<Box mx='auto' style={{ maxWidth, width: '100%', padding: 16 }}>
|
||||
<Stack gap='md'>
|
||||
<Title order={2}>Resources</Title>
|
||||
<TextInput
|
||||
placeholder='Search by Resources ID'
|
||||
value={searchTerm}
|
||||
onChange={(event) => {
|
||||
setSearchTerm(event.currentTarget.value);
|
||||
setPage(1);
|
||||
}}
|
||||
leftSection={<IconSearch size={16} />}
|
||||
data-testid='resources-search-input'
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
<ResourcesTable
|
||||
resourcesList={pagedResources}
|
||||
totalRecords={totalRecordsValue}
|
||||
isFetching={isFetching}
|
||||
isError={isError}
|
||||
error={error}
|
||||
searchTerm={searchTerm}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={(status) => {
|
||||
setSort({
|
||||
column: status.columnAccessor as string,
|
||||
direction: status.direction,
|
||||
});
|
||||
}}
|
||||
onPageChange={setPage}
|
||||
onRecordsPerPageChange={(value) => {
|
||||
setRecordsPerPage(value);
|
||||
setPage(1);
|
||||
}}
|
||||
onResetFilters={() => {
|
||||
setSearchTerm('');
|
||||
setSort({ column: 'resourcesId', direction: 'asc' });
|
||||
setPage(1);
|
||||
}}
|
||||
onRefetch={() => undefined}
|
||||
recordsPerPageOptions={[5, 10, 20]}
|
||||
renderRowExpansion={({ resources }) => <ResourcesTree resources={resources} />}
|
||||
/>
|
||||
</Stack>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
export const WideContainer: Story = {
|
||||
render: () => <ResourcesTableStoryWrapper maxWidth={1280} />,
|
||||
};
|
||||
|
||||
export const MediumContainer: Story = {
|
||||
render: () => <ResourcesTableStoryWrapper maxWidth={960} />,
|
||||
};
|
||||
|
||||
export const NarrowContainer: Story = {
|
||||
render: () => <ResourcesTableStoryWrapper maxWidth={720} />,
|
||||
};
|
||||
|
||||
export const DrawerWidth: Story = {
|
||||
render: () => <ResourcesTableStoryWrapper maxWidth={520} />,
|
||||
};
|
||||
|
||||
export const ErrorState: Story = {
|
||||
render: () => (
|
||||
<ResourcesTableStoryWrapper maxWidth={600} resourcesList={[]} isError error={new Error('Network unreachable')} />
|
||||
),
|
||||
};
|
||||
|
||||
export const EmptyResources: Story = {
|
||||
render: () => (
|
||||
<ResourcesTableStoryWrapper
|
||||
maxWidth={960}
|
||||
resourcesList={[
|
||||
{
|
||||
resourcesId: 'rs-empty-001',
|
||||
version: 1,
|
||||
createTime: 1702000000,
|
||||
updateTime: 1704600000,
|
||||
resources: {},
|
||||
},
|
||||
]}
|
||||
/>
|
||||
),
|
||||
};
|
||||
|
||||
export const LoadingState: Story = {
|
||||
render: () => <ResourcesTableStoryWrapper maxWidth={960} resourcesList={[]} isFetching />,
|
||||
};
|
||||
@@ -0,0 +1,129 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo } from 'react';
|
||||
import { IconAlertCircle, IconChevronRight } from '@tabler/icons-react';
|
||||
import { Box, Group, Stack, Text, Tree, type TreeNodeData } from '@mantine/core';
|
||||
import type { Resources } from '@/types';
|
||||
import { safeStringify } from '@/utils/format';
|
||||
|
||||
function convertToTreeData(obj: any, key: string = 'root', parentPath = ''): TreeNodeData {
|
||||
const isObject = obj !== null && typeof obj === 'object' && !Array.isArray(obj);
|
||||
const isArray = Array.isArray(obj);
|
||||
const currentPath = parentPath ? `${parentPath}.${key}` : key;
|
||||
|
||||
if (isObject) {
|
||||
const children = Object.entries(obj).map(([childKey, childValue]) =>
|
||||
convertToTreeData(childValue, childKey, currentPath),
|
||||
);
|
||||
|
||||
return {
|
||||
value: currentPath,
|
||||
label: (
|
||||
<Group gap={6}>
|
||||
<Text size='sm' fw={500}>
|
||||
{key}
|
||||
</Text>
|
||||
<Text size='xs' c='dimmed'>
|
||||
(Object)
|
||||
</Text>
|
||||
</Group>
|
||||
),
|
||||
children: children.length > 0 ? children : undefined,
|
||||
};
|
||||
}
|
||||
|
||||
if (isArray) {
|
||||
const children = obj.map((item: any, index: number) => convertToTreeData(item, `[${index}]`, currentPath));
|
||||
|
||||
return {
|
||||
value: currentPath,
|
||||
label: (
|
||||
<Group gap={6}>
|
||||
<Text size='sm' fw={500}>
|
||||
{key}
|
||||
</Text>
|
||||
<Text size='xs' c='dimmed'>
|
||||
(Array[
|
||||
{obj.length}
|
||||
])
|
||||
</Text>
|
||||
</Group>
|
||||
),
|
||||
children: children.length > 0 ? children : undefined,
|
||||
};
|
||||
}
|
||||
|
||||
// Primitive value
|
||||
return {
|
||||
value: currentPath,
|
||||
label: (
|
||||
<Group gap={6}>
|
||||
<Text size='sm' fw={500}>
|
||||
{key}:
|
||||
</Text>
|
||||
<Text size='sm' ff='monospace' c='dimmed'>
|
||||
{safeStringify(obj)}
|
||||
</Text>
|
||||
</Group>
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
export type ResourcesTreeProps = {
|
||||
resources: Resources;
|
||||
};
|
||||
|
||||
export function ResourcesTree({ resources }: ResourcesTreeProps) {
|
||||
const resourcesDict = resources.resources ?? {};
|
||||
|
||||
const treeData = useMemo<TreeNodeData[]>(() => {
|
||||
const entries = Object.entries(resourcesDict);
|
||||
|
||||
if (entries.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
return entries.map(([key, value]) => convertToTreeData(value, key));
|
||||
}, [resourcesDict]);
|
||||
|
||||
if (treeData.length === 0) {
|
||||
return (
|
||||
<Stack gap='xs' align='center' py='md'>
|
||||
<IconAlertCircle size={24} color='gray' />
|
||||
<Text size='sm' c='dimmed'>
|
||||
No resources found
|
||||
</Text>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<Box p='md' style={{ backgroundColor: 'var(--mantine-color-default-hover)' }}>
|
||||
<Tree
|
||||
data={treeData}
|
||||
levelOffset={20}
|
||||
expandOnClick
|
||||
selectOnClick
|
||||
renderNode={({ node, expanded, hasChildren, elementProps }) => (
|
||||
<Group gap={4} {...elementProps}>
|
||||
{hasChildren && (
|
||||
<Box
|
||||
style={{
|
||||
minWidth: 14,
|
||||
display: 'flex',
|
||||
alignItems: 'center',
|
||||
justifyContent: 'center',
|
||||
transform: expanded ? 'rotate(90deg)' : 'rotate(0deg)',
|
||||
transition: 'transform 150ms ease',
|
||||
}}
|
||||
>
|
||||
<IconChevronRight size={14} />
|
||||
</Box>
|
||||
)}
|
||||
{node.label}
|
||||
</Group>
|
||||
)}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { Box, Stack, Title } from '@mantine/core';
|
||||
import type { Resources } from '@/types';
|
||||
import { ResourcesTree } from './ResourcesTree.component';
|
||||
|
||||
const meta: Meta<typeof ResourcesTree> = {
|
||||
title: 'Components/ResourcesTree',
|
||||
component: ResourcesTree,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof ResourcesTree>;
|
||||
|
||||
const simpleResources: Resources = {
|
||||
resourcesId: 'rs-simple-001',
|
||||
version: 1,
|
||||
createTime: 1704067200,
|
||||
updateTime: 1706668800,
|
||||
resources: {
|
||||
apiKey: { value: 'sk-test-key-123', type: 'secret' },
|
||||
maxRetries: { value: 3, description: 'Maximum retry attempts' },
|
||||
timeout: { value: 30000, unit: 'ms' },
|
||||
enabled: { value: true },
|
||||
},
|
||||
};
|
||||
|
||||
const nestedResources: Resources = {
|
||||
resourcesId: 'rs-nested-001',
|
||||
version: 2,
|
||||
createTime: 1709251200,
|
||||
updateTime: 1711856800,
|
||||
resources: {
|
||||
model: {
|
||||
name: 'gpt-4',
|
||||
version: '2024-01-01',
|
||||
temperature: 0.7,
|
||||
maxTokens: 2048,
|
||||
topP: 0.9,
|
||||
},
|
||||
database: {
|
||||
host: 'db.example.com',
|
||||
port: 5432,
|
||||
name: 'production',
|
||||
pool: {
|
||||
min: 2,
|
||||
max: 10,
|
||||
idle: 30000,
|
||||
},
|
||||
},
|
||||
cache: {
|
||||
type: 'redis',
|
||||
host: 'cache.example.com',
|
||||
port: 6379,
|
||||
ttl: 3600,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const arrayResources: Resources = {
|
||||
resourcesId: 'rs-array-001',
|
||||
version: 3,
|
||||
createTime: 1712217600,
|
||||
updateTime: 1714823200,
|
||||
resources: {
|
||||
compute: {
|
||||
instances: [
|
||||
{ id: 'i-001', type: 't3.large', zone: 'us-east-1a', status: 'running' },
|
||||
{ id: 'i-002', type: 't3.large', zone: 'us-east-1b', status: 'running' },
|
||||
{ id: 'i-003', type: 't3.xlarge', zone: 'us-east-1c', status: 'stopped' },
|
||||
],
|
||||
autoScaling: {
|
||||
min: 2,
|
||||
max: 10,
|
||||
targetCpu: 70,
|
||||
},
|
||||
},
|
||||
tags: ['production', 'ml-training', 'auto-scale'],
|
||||
ports: [80, 443, 8080],
|
||||
},
|
||||
};
|
||||
|
||||
const complexResources: Resources = {
|
||||
resourcesId: 'rs-complex-001',
|
||||
version: 4,
|
||||
createTime: 1706745600,
|
||||
updateTime: 1710000000,
|
||||
resources: {
|
||||
model: {
|
||||
name: 'claude-3-opus',
|
||||
version: '2024-02-01',
|
||||
temperature: 0.5,
|
||||
maxTokens: 4096,
|
||||
providers: [
|
||||
{ name: 'anthropic', priority: 1, enabled: true },
|
||||
{ name: 'aws-bedrock', priority: 2, enabled: false },
|
||||
],
|
||||
},
|
||||
storage: {
|
||||
type: 's3',
|
||||
bucket: 'training-data',
|
||||
region: 'us-east-1',
|
||||
credentials: {
|
||||
accessKeyId: 'AKIA***',
|
||||
encrypted: true,
|
||||
},
|
||||
lifecycle: {
|
||||
transitionToIA: 30,
|
||||
transitionToGlacier: 90,
|
||||
expiration: 365,
|
||||
},
|
||||
},
|
||||
monitoring: {
|
||||
enabled: true,
|
||||
interval: 60,
|
||||
metrics: ['cpu', 'memory', 'disk', 'network'],
|
||||
alerts: {
|
||||
email: 'ops@example.com',
|
||||
slack: '#alerts',
|
||||
pagerduty: true,
|
||||
thresholds: {
|
||||
cpu: { warning: 70, critical: 90 },
|
||||
memory: { warning: 80, critical: 95 },
|
||||
disk: { warning: 75, critical: 90 },
|
||||
},
|
||||
},
|
||||
},
|
||||
apiKeys: {
|
||||
openai: 'sk-***',
|
||||
anthropic: 'sk-ant-***',
|
||||
replicate: 'r8-***',
|
||||
},
|
||||
rateLimits: {
|
||||
requestsPerMinute: 100,
|
||||
tokensPerDay: 1000000,
|
||||
concurrent: 5,
|
||||
burstMultiplier: 1.5,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const emptyResources: Resources = {
|
||||
resourcesId: 'rs-empty-001',
|
||||
version: 1,
|
||||
createTime: 1702000000,
|
||||
updateTime: 1704600000,
|
||||
resources: {},
|
||||
};
|
||||
|
||||
type WrapperProps = {
|
||||
resources: Resources;
|
||||
maxWidth?: number;
|
||||
};
|
||||
|
||||
function ResourcesTreeStoryWrapper({ resources, maxWidth = 800 }: WrapperProps) {
|
||||
return (
|
||||
<Box mx='auto' style={{ maxWidth, width: '100%', padding: 16 }}>
|
||||
<Stack gap='md'>
|
||||
<Title order={2}>
|
||||
Resources:
|
||||
{resources.resourcesId}
|
||||
</Title>
|
||||
<ResourcesTree resources={resources} />
|
||||
</Stack>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
export const SimpleValues: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={simpleResources} />,
|
||||
};
|
||||
|
||||
export const NestedObjects: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={nestedResources} />,
|
||||
};
|
||||
|
||||
export const WithArrays: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={arrayResources} />,
|
||||
};
|
||||
|
||||
export const ComplexStructure: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={complexResources} maxWidth={1000} />,
|
||||
};
|
||||
|
||||
export const EmptyResources: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={emptyResources} />,
|
||||
};
|
||||
|
||||
export const NarrowContainer: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={complexResources} maxWidth={500} />,
|
||||
};
|
||||
|
||||
export const WideContainer: Story = {
|
||||
render: () => <ResourcesTreeStoryWrapper resources={complexResources} maxWidth={1400} />,
|
||||
};
|
||||
@@ -0,0 +1,787 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo, useState, type ReactNode, type SetStateAction } from 'react';
|
||||
import {
|
||||
IconAlertCircle,
|
||||
IconCheck,
|
||||
IconCopy,
|
||||
IconFileDescription,
|
||||
IconRefresh,
|
||||
IconReload,
|
||||
IconTimeline,
|
||||
} from '@tabler/icons-react';
|
||||
import { DataTable, type DataTableColumn, type DataTableSortStatus } from 'mantine-datatable';
|
||||
import {
|
||||
ActionIcon,
|
||||
Alert,
|
||||
Badge,
|
||||
Box,
|
||||
Button,
|
||||
CopyButton,
|
||||
Group,
|
||||
MultiSelect,
|
||||
Stack,
|
||||
Text,
|
||||
Tooltip,
|
||||
} from '@mantine/core';
|
||||
import { useElementSize, useViewportSize } from '@mantine/hooks';
|
||||
import {
|
||||
type Attempt,
|
||||
type AttemptStatus,
|
||||
type Rollout,
|
||||
type RolloutMode,
|
||||
type RolloutsSortState,
|
||||
type RolloutStatus,
|
||||
} from '@/features/rollouts';
|
||||
import { getLayoutAwareWidth } from '@/layouts/helper';
|
||||
import {
|
||||
clampToNow,
|
||||
formatDateTime,
|
||||
formatDuration,
|
||||
formatRelativeTime,
|
||||
formatStatusLabel,
|
||||
safeStringify,
|
||||
toTimestamp,
|
||||
} from '@/utils/format';
|
||||
import { createResponsiveColumns, type ColumnVisibilityConfig } from '@/utils/table';
|
||||
|
||||
const ROLLOUT_STATUS_OPTIONS: RolloutStatus[] = [
|
||||
'queuing',
|
||||
'preparing',
|
||||
'running',
|
||||
'failed',
|
||||
'succeeded',
|
||||
'cancelled',
|
||||
'requeuing',
|
||||
];
|
||||
|
||||
const ATTEMPT_STATUS_COLORS: Record<AttemptStatus, string> = {
|
||||
failed: 'red',
|
||||
preparing: 'violet',
|
||||
running: 'blue',
|
||||
succeeded: 'teal',
|
||||
timeout: 'orange',
|
||||
unresponsive: 'orange',
|
||||
};
|
||||
|
||||
const ROLLOUT_STATUS_COLORS: Record<RolloutStatus, string> = {
|
||||
cancelled: 'gray',
|
||||
failed: 'red',
|
||||
preparing: 'violet',
|
||||
queuing: 'gray',
|
||||
requeuing: 'gray',
|
||||
running: 'blue',
|
||||
succeeded: 'teal',
|
||||
};
|
||||
|
||||
const ROLLOUT_MODE_OPTIONS: RolloutMode[] = ['train', 'val', 'test'];
|
||||
|
||||
const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
rolloutId: { fixedWidth: 12.5, priority: 0 },
|
||||
actionsPlaceholder: { fixedWidth: 6.5, priority: 0 },
|
||||
inputText: { minWidth: 14, priority: 1 },
|
||||
statusValue: { fixedWidth: 10, priority: 1 },
|
||||
startTimestamp: { fixedWidth: 12, priority: 2 },
|
||||
durationSeconds: { fixedWidth: 10, priority: 2 },
|
||||
attemptId: { fixedWidth: 12, priority: 3 },
|
||||
resourcesId: { fixedWidth: 10, priority: 3 },
|
||||
mode: { fixedWidth: 8, priority: 3 },
|
||||
lastHeartbeatTimestamp: { fixedWidth: 10, priority: 3 },
|
||||
workerId: { fixedWidth: 10, priority: 3 },
|
||||
};
|
||||
|
||||
export type RolloutTableRecord = Rollout & {
|
||||
attemptId: string | null;
|
||||
attemptSequence: number | null;
|
||||
isNested: boolean;
|
||||
canExpand: boolean;
|
||||
inputText: string;
|
||||
attemptStatus?: AttemptStatus;
|
||||
statusValue: string;
|
||||
startTimestamp: number | null;
|
||||
durationSeconds: number | null;
|
||||
lastHeartbeatTimestamp: number | null;
|
||||
workerId: string | null;
|
||||
actionsPlaceholder?: null;
|
||||
};
|
||||
|
||||
function selectHeartbeatTimestamp(attempt?: Attempt | null): number | null {
|
||||
if (!attempt || attempt.lastHeartbeatTime == null || Number.isNaN(attempt.lastHeartbeatTime)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return attempt.lastHeartbeatTime;
|
||||
}
|
||||
|
||||
export function buildRolloutRecord(rollout: Rollout): RolloutTableRecord {
|
||||
const latestAttempt = rollout.attempt;
|
||||
const inputValue =
|
||||
rollout.input === null || typeof rollout.input === 'undefined'
|
||||
? '—'
|
||||
: typeof rollout.input === 'string'
|
||||
? rollout.input
|
||||
: safeStringify(rollout.input);
|
||||
const startTimestamp = toTimestamp(latestAttempt?.startTime ?? rollout.startTime);
|
||||
const endTimestamp = toTimestamp(latestAttempt?.endTime ?? rollout.endTime);
|
||||
const durationSeconds = clampToNow(startTimestamp, endTimestamp);
|
||||
const attemptStatus = latestAttempt?.status;
|
||||
const sequenceId = latestAttempt?.sequenceId;
|
||||
const statusValue =
|
||||
attemptStatus && attemptStatus !== rollout.status ? `${rollout.status}-${attemptStatus}` : rollout.status;
|
||||
|
||||
return {
|
||||
...rollout,
|
||||
attempt: latestAttempt ?? null,
|
||||
attemptId: latestAttempt?.attemptId ?? null,
|
||||
attemptSequence: latestAttempt?.sequenceId ?? null,
|
||||
isNested: false,
|
||||
canExpand: Boolean(sequenceId && sequenceId > 1),
|
||||
inputText: inputValue,
|
||||
attemptStatus,
|
||||
statusValue,
|
||||
startTimestamp,
|
||||
durationSeconds,
|
||||
lastHeartbeatTimestamp: rollout.attempt?.lastHeartbeatTime ?? null,
|
||||
workerId: latestAttempt?.workerId ?? null,
|
||||
actionsPlaceholder: null,
|
||||
};
|
||||
}
|
||||
|
||||
function buildAttemptRecord(rollout: Rollout, attempt: Attempt): RolloutTableRecord {
|
||||
const inputValue =
|
||||
rollout.input === null || typeof rollout.input === 'undefined'
|
||||
? '—'
|
||||
: typeof rollout.input === 'string'
|
||||
? rollout.input
|
||||
: safeStringify(rollout.input);
|
||||
const startTimestamp = toTimestamp(attempt.startTime ?? rollout.startTime);
|
||||
const endTimestamp = toTimestamp(attempt.endTime);
|
||||
const durationSeconds = clampToNow(startTimestamp, endTimestamp);
|
||||
const lastHeartbeatTimestamp = selectHeartbeatTimestamp(attempt);
|
||||
|
||||
return {
|
||||
...rollout,
|
||||
attempt,
|
||||
attemptId: attempt.attemptId,
|
||||
attemptSequence: attempt.sequenceId,
|
||||
isNested: true,
|
||||
canExpand: false,
|
||||
inputText: inputValue,
|
||||
attemptStatus: attempt.status,
|
||||
statusValue: attempt.status,
|
||||
startTimestamp,
|
||||
durationSeconds,
|
||||
lastHeartbeatTimestamp,
|
||||
workerId: attempt.workerId ?? null,
|
||||
actionsPlaceholder: null,
|
||||
};
|
||||
}
|
||||
|
||||
function getStatusBadge(status: string, kind: 'rollout' | 'attempt') {
|
||||
const color =
|
||||
kind === 'rollout'
|
||||
? (ROLLOUT_STATUS_COLORS[status as RolloutStatus] ?? 'gray')
|
||||
: (ATTEMPT_STATUS_COLORS[status as AttemptStatus] ?? 'gray');
|
||||
|
||||
return (
|
||||
<Badge size='sm' variant='light' color={color}>
|
||||
{formatStatusLabel(status)}
|
||||
</Badge>
|
||||
);
|
||||
}
|
||||
|
||||
type RolloutColumnsOptions = {
|
||||
statusFilters: RolloutStatus[];
|
||||
onStatusFilterChange: (values: RolloutStatus[]) => void;
|
||||
onStatusFilterReset: () => void;
|
||||
modeFilters: RolloutMode[];
|
||||
onModeFilterChange: (values: RolloutMode[]) => void;
|
||||
onModeFilterReset: () => void;
|
||||
onViewRawJson?: (record: RolloutTableRecord) => void;
|
||||
onViewTraces?: (record: RolloutTableRecord) => void;
|
||||
};
|
||||
|
||||
function createRolloutColumns({
|
||||
statusFilters,
|
||||
onStatusFilterChange,
|
||||
onStatusFilterReset,
|
||||
modeFilters,
|
||||
onModeFilterChange,
|
||||
onModeFilterReset,
|
||||
onViewRawJson,
|
||||
onViewTraces,
|
||||
}: RolloutColumnsOptions): DataTableColumn<RolloutTableRecord>[] {
|
||||
const statusOptions = ROLLOUT_STATUS_OPTIONS.map((status) => ({
|
||||
value: status,
|
||||
label: formatStatusLabel(status),
|
||||
}));
|
||||
const modeOptions = ROLLOUT_MODE_OPTIONS.map((mode) => ({
|
||||
value: mode,
|
||||
label: formatStatusLabel(mode),
|
||||
}));
|
||||
|
||||
return [
|
||||
{
|
||||
accessor: 'rolloutId',
|
||||
title: 'Rollout',
|
||||
sortable: true,
|
||||
render: ({ rolloutId }) => (
|
||||
<Group gap={2}>
|
||||
<Text fw={500} size='sm'>
|
||||
{rolloutId}
|
||||
</Text>
|
||||
<CopyButton value={rolloutId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy rollout ID ${rolloutId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'attemptId',
|
||||
title: 'Attempt',
|
||||
sortable: true,
|
||||
render: ({ attemptId, attemptSequence, isNested }) => (
|
||||
<Group gap={2}>
|
||||
<Text size='sm' c={attemptId ? undefined : 'dimmed'}>
|
||||
{attemptId ?? '—'}
|
||||
</Text>
|
||||
{attemptId && (
|
||||
<CopyButton value={attemptId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy attempt ID ${attemptId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
)}
|
||||
{attemptSequence && (isNested || attemptSequence > 1) && (
|
||||
<Badge leftSection={<IconReload size={12} />} pl={6} pr={6}>
|
||||
{attemptSequence}
|
||||
</Badge>
|
||||
)}
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'inputText',
|
||||
title: 'Input',
|
||||
render: ({ inputText }) => (
|
||||
<Text
|
||||
size='sm'
|
||||
ff='monospace'
|
||||
c='dimmed'
|
||||
lineClamp={1}
|
||||
title={inputText}
|
||||
style={{ width: '100%', wordBreak: 'break-all', overflow: 'hidden' }}
|
||||
>
|
||||
{inputText}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'statusValue',
|
||||
title: 'Status',
|
||||
sortable: true,
|
||||
filter: ({ close }) => (
|
||||
<Stack gap='xs'>
|
||||
<MultiSelect
|
||||
label='Status'
|
||||
description='Filter rollouts by status'
|
||||
data={statusOptions}
|
||||
value={statusFilters}
|
||||
placeholder='Select statuses...'
|
||||
searchable
|
||||
clearable
|
||||
comboboxProps={{ withinPortal: false }}
|
||||
onChange={(values) => onStatusFilterChange(values as RolloutStatus[])}
|
||||
/>
|
||||
<Button
|
||||
variant='light'
|
||||
size='xs'
|
||||
onClick={() => {
|
||||
onStatusFilterReset();
|
||||
close();
|
||||
}}
|
||||
disabled={statusFilters.length === 0}
|
||||
>
|
||||
Clear
|
||||
</Button>
|
||||
</Stack>
|
||||
),
|
||||
filtering: statusFilters.length > 0,
|
||||
render: ({ status, attemptStatus, isNested }) => {
|
||||
if (isNested) {
|
||||
return <Group gap={4}>{getStatusBadge(attemptStatus ?? 'unknown', 'attempt')}</Group>;
|
||||
}
|
||||
|
||||
if (attemptStatus && attemptStatus !== status) {
|
||||
return (
|
||||
<Group gap={4}>
|
||||
{getStatusBadge(status, 'rollout')}
|
||||
<Text size='sm' c='dimmed'>
|
||||
—
|
||||
</Text>
|
||||
{getStatusBadge(attemptStatus, 'attempt')}
|
||||
</Group>
|
||||
);
|
||||
}
|
||||
|
||||
return getStatusBadge(status, 'rollout');
|
||||
},
|
||||
},
|
||||
{
|
||||
accessor: 'resourcesId',
|
||||
title: 'Resources',
|
||||
sortable: true,
|
||||
render: ({ resourcesId }) => (
|
||||
<Text size='sm' c={resourcesId ? undefined : 'dimmed'}>
|
||||
{resourcesId ?? '—'}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'mode',
|
||||
title: 'Mode',
|
||||
sortable: true,
|
||||
filter: ({ close }) => (
|
||||
<Stack gap='xs'>
|
||||
<MultiSelect
|
||||
label='Mode'
|
||||
description='Filter rollouts by mode'
|
||||
data={modeOptions}
|
||||
value={modeFilters}
|
||||
placeholder='Select modes...'
|
||||
searchable
|
||||
clearable
|
||||
comboboxProps={{ withinPortal: false }}
|
||||
onChange={(values) => onModeFilterChange(values as RolloutMode[])}
|
||||
/>
|
||||
<Button
|
||||
variant='light'
|
||||
size='xs'
|
||||
onClick={() => {
|
||||
onModeFilterReset();
|
||||
close();
|
||||
}}
|
||||
disabled={modeFilters.length === 0}
|
||||
>
|
||||
Clear
|
||||
</Button>
|
||||
</Stack>
|
||||
),
|
||||
filtering: modeFilters.length > 0,
|
||||
render: ({ mode }) => (
|
||||
<Text size='sm' c={mode ? undefined : 'dimmed'}>
|
||||
{mode ?? '—'}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'startTimestamp',
|
||||
title: 'Start Time',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ startTimestamp }) => <Text size='sm'>{formatDateTime(startTimestamp)}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'durationSeconds',
|
||||
title: 'Duration',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ durationSeconds }) => <Text size='sm'>{formatDuration(durationSeconds)}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'lastHeartbeatTimestamp',
|
||||
title: 'Last Heartbeat',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ lastHeartbeatTimestamp, attempt, isNested }) => {
|
||||
if (!attempt && isNested) {
|
||||
return (
|
||||
<Text size='sm' c='dimmed'>
|
||||
—
|
||||
</Text>
|
||||
);
|
||||
}
|
||||
return <Text size='sm'>{formatRelativeTime(lastHeartbeatTimestamp)}</Text>;
|
||||
},
|
||||
},
|
||||
{
|
||||
accessor: 'workerId',
|
||||
title: 'Worker',
|
||||
sortable: true,
|
||||
render: ({ workerId }) => (
|
||||
<Text size='sm' c={workerId ? undefined : 'dimmed'}>
|
||||
{workerId ?? '—'}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'actionsPlaceholder',
|
||||
title: 'Actions',
|
||||
render: (record) => (
|
||||
<Group gap={4}>
|
||||
<Tooltip label='View raw JSON' withArrow disabled={!onViewRawJson}>
|
||||
<ActionIcon
|
||||
aria-label='View raw JSON'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onViewRawJson?.(record);
|
||||
}}
|
||||
>
|
||||
<IconFileDescription size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
<Tooltip label='View traces' withArrow disabled={!onViewTraces}>
|
||||
<ActionIcon
|
||||
aria-label='View traces'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onViewTraces?.(record);
|
||||
}}
|
||||
>
|
||||
<IconTimeline size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
type RowExpansionRenderer = (context: {
|
||||
rollout: Rollout;
|
||||
columns: DataTableColumn<RolloutTableRecord>[];
|
||||
}) => ReactNode;
|
||||
|
||||
export type RolloutTableProps = {
|
||||
rollouts: Rollout[] | undefined;
|
||||
totalRecords: number;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
error: unknown;
|
||||
searchTerm: string;
|
||||
statusFilters: RolloutStatus[];
|
||||
modeFilters: RolloutMode[];
|
||||
sort: RolloutsSortState;
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
onStatusFilterChange: (values: RolloutStatus[]) => void;
|
||||
onStatusFilterReset: () => void;
|
||||
onModeFilterChange: (values: RolloutMode[]) => void;
|
||||
onModeFilterReset: () => void;
|
||||
onSortStatusChange: (status: DataTableSortStatus<RolloutTableRecord>) => void;
|
||||
onPageChange: (page: number) => void;
|
||||
onRecordsPerPageChange: (value: number) => void;
|
||||
onResetFilters: () => void;
|
||||
onRefetch: () => void;
|
||||
onViewRawJson?: (record: RolloutTableRecord) => void;
|
||||
onViewTraces?: (record: RolloutTableRecord) => void;
|
||||
recordsPerPageOptions?: number[];
|
||||
renderRowExpansion?: RowExpansionRenderer;
|
||||
};
|
||||
|
||||
export function RolloutTable({
|
||||
rollouts,
|
||||
totalRecords,
|
||||
isFetching,
|
||||
isError,
|
||||
error,
|
||||
searchTerm,
|
||||
statusFilters,
|
||||
modeFilters,
|
||||
sort,
|
||||
page,
|
||||
recordsPerPage,
|
||||
onStatusFilterChange,
|
||||
onStatusFilterReset,
|
||||
onModeFilterChange,
|
||||
onModeFilterReset,
|
||||
onSortStatusChange,
|
||||
onPageChange,
|
||||
onRecordsPerPageChange,
|
||||
onResetFilters,
|
||||
onRefetch,
|
||||
onViewRawJson,
|
||||
onViewTraces,
|
||||
recordsPerPageOptions = DEFAULT_RECORDS_PER_PAGE_OPTIONS,
|
||||
renderRowExpansion,
|
||||
}: RolloutTableProps) {
|
||||
const [expandedRecordIds, setExpandedRecordIds] = useState<string[]>([]);
|
||||
const { ref: tableContainerRef, width: containerWidth } = useElementSize();
|
||||
const { width: viewportWidth } = useViewportSize();
|
||||
|
||||
const layoutAwareContainerWidth = useMemo(() => {
|
||||
return getLayoutAwareWidth(containerWidth, viewportWidth);
|
||||
}, [containerWidth, viewportWidth]);
|
||||
|
||||
const rolloutRecords = useMemo<RolloutTableRecord[]>(() => {
|
||||
if (!rollouts) {
|
||||
return [];
|
||||
}
|
||||
return rollouts.map((rolloutItem) => buildRolloutRecord(rolloutItem));
|
||||
}, [rollouts]);
|
||||
|
||||
const columns = useMemo(
|
||||
() =>
|
||||
createRolloutColumns({
|
||||
statusFilters,
|
||||
onStatusFilterChange,
|
||||
onStatusFilterReset,
|
||||
modeFilters,
|
||||
onModeFilterChange,
|
||||
onModeFilterReset,
|
||||
onViewRawJson,
|
||||
onViewTraces,
|
||||
}),
|
||||
[
|
||||
statusFilters,
|
||||
onStatusFilterChange,
|
||||
onStatusFilterReset,
|
||||
modeFilters,
|
||||
onModeFilterChange,
|
||||
onModeFilterReset,
|
||||
onViewRawJson,
|
||||
onViewTraces,
|
||||
],
|
||||
);
|
||||
|
||||
const responsiveColumns = useMemo(
|
||||
() => createResponsiveColumns(columns, layoutAwareContainerWidth, COLUMN_VISIBILITY),
|
||||
[columns, layoutAwareContainerWidth],
|
||||
);
|
||||
|
||||
const totalPages = useMemo(
|
||||
() => Math.max(1, Math.ceil(Math.max(0, totalRecords) / Math.max(1, recordsPerPage))),
|
||||
[recordsPerPage, totalRecords],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (page > totalPages) {
|
||||
onPageChange(totalPages);
|
||||
}
|
||||
}, [onPageChange, page, totalPages]);
|
||||
|
||||
useEffect(() => {
|
||||
setExpandedRecordIds((current) =>
|
||||
current.filter((id) => rolloutRecords.some((record) => record.rolloutId === id && record.canExpand)),
|
||||
);
|
||||
}, [rolloutRecords]);
|
||||
|
||||
const hasActiveFilters = searchTerm.trim().length > 0 || statusFilters.length > 0 || modeFilters.length > 0;
|
||||
|
||||
const sortStatus: DataTableSortStatus<RolloutTableRecord> = {
|
||||
columnAccessor: sort.column,
|
||||
direction: sort.direction,
|
||||
};
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<RolloutTableRecord>) => {
|
||||
onSortStatusChange(status);
|
||||
},
|
||||
[onSortStatusChange],
|
||||
);
|
||||
|
||||
const errorMessage =
|
||||
isError && error && typeof error === 'object' && 'status' in (error as Record<string, unknown>)
|
||||
? `Rollouts are temporarily unavailable (status: ${String((error as Record<string, unknown>).status)}).`
|
||||
: 'Rollouts are temporarily unavailable.';
|
||||
|
||||
const emptyState = (
|
||||
<Stack gap='sm' align='center' py='lg'>
|
||||
{isError ? (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
{errorMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Use the retry button to try again, or adjust the filters to broaden the results.
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' color='gray' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Retry
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
No rollouts found
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
{hasActiveFilters
|
||||
? 'Try adjusting the search or filters to see more results.'
|
||||
: 'Try refreshing to fetch the latest rollouts.'}
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Refresh
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
|
||||
return (
|
||||
<Box ref={tableContainerRef} data-testid='rollouts-table-container'>
|
||||
<DataTable<RolloutTableRecord>
|
||||
classNames={{ root: 'rollouts-table' }}
|
||||
withTableBorder
|
||||
withColumnBorders
|
||||
highlightOnHover
|
||||
verticalAlign='center'
|
||||
minHeight={rolloutRecords.length === 0 ? 500 : undefined}
|
||||
idAccessor='rolloutId'
|
||||
records={rolloutRecords}
|
||||
columns={responsiveColumns}
|
||||
totalRecords={totalRecords}
|
||||
recordsPerPage={recordsPerPage}
|
||||
page={page}
|
||||
onPageChange={onPageChange}
|
||||
onRecordsPerPageChange={onRecordsPerPageChange}
|
||||
recordsPerPageOptions={recordsPerPageOptions}
|
||||
sortStatus={sortStatus}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
emptyState={rolloutRecords.length === 0 ? emptyState : undefined}
|
||||
rowExpansion={
|
||||
renderRowExpansion
|
||||
? {
|
||||
allowMultiple: true,
|
||||
expandable: ({ record }) => record.canExpand,
|
||||
expanded: {
|
||||
recordIds: expandedRecordIds,
|
||||
onRecordIdsChange: (nextRecordIds: SetStateAction<string[]>) => {
|
||||
setExpandedRecordIds((previous) => {
|
||||
const resolved =
|
||||
typeof nextRecordIds === 'function'
|
||||
? nextRecordIds(previous)
|
||||
: ((nextRecordIds ?? []) as (string | number)[]);
|
||||
return resolved
|
||||
.map(String)
|
||||
.filter((id) =>
|
||||
rolloutRecords.some((tableRecord) => tableRecord.rolloutId === id && tableRecord.canExpand),
|
||||
);
|
||||
});
|
||||
},
|
||||
},
|
||||
content: ({ record }) => renderRowExpansion({ rollout: record, columns: responsiveColumns }),
|
||||
}
|
||||
: undefined
|
||||
}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
export type RolloutAttemptsTableProps = {
|
||||
rollout: Rollout;
|
||||
attempts: Attempt[] | undefined;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
onRetry: () => void;
|
||||
columns: DataTableColumn<RolloutTableRecord>[];
|
||||
};
|
||||
|
||||
export function RolloutAttemptsTable({
|
||||
rollout,
|
||||
attempts,
|
||||
isFetching,
|
||||
isError,
|
||||
onRetry,
|
||||
columns,
|
||||
}: RolloutAttemptsTableProps) {
|
||||
const attemptRecords = useMemo<RolloutTableRecord[]>(() => {
|
||||
if (!attempts) {
|
||||
return [];
|
||||
}
|
||||
return attempts
|
||||
.map((attempt) => buildAttemptRecord(rollout, attempt))
|
||||
.sort((a, b) => (b.attemptSequence ?? 0) - (a.attemptSequence ?? 0))
|
||||
.filter((record) => record.attemptSequence !== rollout.attempt?.sequenceId);
|
||||
}, [attempts, rollout]);
|
||||
|
||||
if (isError && !attemptRecords.length) {
|
||||
return (
|
||||
<Alert color='red' variant='light' icon={<IconAlertCircle size={16} />}>
|
||||
<Stack gap='xs'>
|
||||
<Text size='sm'>Unable to load attempts for this rollout.</Text>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRetry}>
|
||||
Retry
|
||||
</Button>
|
||||
</Stack>
|
||||
</Alert>
|
||||
);
|
||||
}
|
||||
|
||||
const emptyState = (
|
||||
<Stack gap='xs' align='center' py='md'>
|
||||
<Text size='sm' c='dimmed'>
|
||||
No attempts found for this rollout.
|
||||
</Text>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRetry}>
|
||||
Refresh
|
||||
</Button>
|
||||
</Stack>
|
||||
);
|
||||
|
||||
return (
|
||||
<DataTable<RolloutTableRecord>
|
||||
classNames={{ root: 'rollouts-table rollouts-table--nested' }}
|
||||
withColumnBorders
|
||||
noHeader
|
||||
minHeight={0}
|
||||
idAccessor='attemptId'
|
||||
verticalAlign='center'
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
records={attemptRecords}
|
||||
columns={columns}
|
||||
emptyState={attemptRecords.length === 0 ? emptyState : undefined}
|
||||
/>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,307 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo, useState } from 'react';
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import { Box, Stack, TextInput, Title } from '@mantine/core';
|
||||
import type { RolloutsSortState } from '@/features/rollouts';
|
||||
import type { Rollout, RolloutMode, RolloutStatus } from '@/types';
|
||||
import { compareRecords } from '@/utils/table';
|
||||
import { STORY_DATE_NOW_SECONDS } from '../../.storybook/constants';
|
||||
import { buildRolloutRecord, RolloutTable, type RolloutTableRecord } from './RolloutTable.component';
|
||||
|
||||
const meta: Meta<typeof RolloutTable> = {
|
||||
title: 'Components/RolloutTable',
|
||||
component: RolloutTable,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof RolloutTable>;
|
||||
|
||||
const now = STORY_DATE_NOW_SECONDS;
|
||||
|
||||
const sampleRollouts: Rollout[] = [
|
||||
{
|
||||
rolloutId: 'ro-story-001',
|
||||
input: { task: 'Generate onboarding summary' },
|
||||
startTime: now - 3200,
|
||||
endTime: null,
|
||||
mode: 'train',
|
||||
resourcesId: 'rs-story-001',
|
||||
status: 'running',
|
||||
config: { retries: 1 },
|
||||
metadata: { owner: 'alice' },
|
||||
attempt: {
|
||||
rolloutId: 'ro-story-001',
|
||||
attemptId: 'at-story-010',
|
||||
sequenceId: 1,
|
||||
startTime: now - 3200,
|
||||
endTime: null,
|
||||
status: 'running',
|
||||
workerId: 'worker-east',
|
||||
lastHeartbeatTime: now - 45,
|
||||
metadata: { info: 'Worker is processing' },
|
||||
},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-story-002',
|
||||
input: { task: 'Classify feedback tickets' },
|
||||
startTime: now - 7200,
|
||||
endTime: now - 5400,
|
||||
mode: 'val',
|
||||
resourcesId: 'rs-story-002',
|
||||
status: 'succeeded',
|
||||
config: { retries: 2 },
|
||||
metadata: { owner: 'bob' },
|
||||
attempt: {
|
||||
rolloutId: 'ro-story-002',
|
||||
attemptId: 'at-story-011',
|
||||
sequenceId: 2,
|
||||
startTime: now - 6200,
|
||||
endTime: now - 5400,
|
||||
status: 'succeeded',
|
||||
workerId: 'worker-north',
|
||||
lastHeartbeatTime: now - 5400,
|
||||
metadata: { previousAttempt: 'at-story-010' },
|
||||
},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-story-003',
|
||||
input: { task: 'Analyze experiment results' },
|
||||
startTime: now - 10800,
|
||||
endTime: now - 9600,
|
||||
mode: 'test',
|
||||
resourcesId: 'rs-story-003',
|
||||
status: 'failed',
|
||||
config: { retries: 1 },
|
||||
metadata: { owner: 'carol' },
|
||||
attempt: {
|
||||
rolloutId: 'ro-story-003',
|
||||
attemptId: 'at-story-012',
|
||||
sequenceId: 3,
|
||||
startTime: now - 10200,
|
||||
endTime: now - 9600,
|
||||
status: 'failed',
|
||||
workerId: 'worker-west',
|
||||
lastHeartbeatTime: now - 9600,
|
||||
metadata: { reason: 'Timeout' },
|
||||
},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-story-004',
|
||||
input: { task: 'Evaluate prompt variants' },
|
||||
startTime: now - 3600,
|
||||
endTime: null,
|
||||
mode: 'train',
|
||||
resourcesId: null,
|
||||
status: 'preparing',
|
||||
config: { retries: 0 },
|
||||
metadata: { owner: 'dave' },
|
||||
attempt: null,
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-story-005',
|
||||
input: { task: 'Generate quick answers' },
|
||||
startTime: now - 1800,
|
||||
endTime: null,
|
||||
mode: 'val',
|
||||
resourcesId: 'rs-story-004',
|
||||
status: 'running',
|
||||
config: { retries: 0 },
|
||||
metadata: { owner: 'eva' },
|
||||
attempt: {
|
||||
rolloutId: 'ro-story-005',
|
||||
attemptId: 'at-story-013',
|
||||
sequenceId: 1,
|
||||
startTime: now - 1800,
|
||||
endTime: null,
|
||||
status: 'running',
|
||||
workerId: null,
|
||||
lastHeartbeatTime: now - 75,
|
||||
metadata: null,
|
||||
},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-story-006',
|
||||
input: { task: 'Compile release notes' },
|
||||
startTime: now - 9600,
|
||||
endTime: now - 9000,
|
||||
mode: null,
|
||||
resourcesId: 'rs-story-005',
|
||||
status: 'cancelled',
|
||||
config: { retries: 3 },
|
||||
metadata: null,
|
||||
attempt: {
|
||||
rolloutId: 'ro-story-006',
|
||||
attemptId: 'at-story-014',
|
||||
sequenceId: 1,
|
||||
startTime: now - 9600,
|
||||
endTime: now - 9000,
|
||||
status: 'timeout',
|
||||
workerId: 'worker-south',
|
||||
lastHeartbeatTime: now - 9000,
|
||||
metadata: { info: 'Cancelled by operator' },
|
||||
},
|
||||
},
|
||||
];
|
||||
|
||||
type WrapperProps = {
|
||||
maxWidth: number;
|
||||
rollouts?: Rollout[] | undefined;
|
||||
isFetching?: boolean;
|
||||
isError?: boolean;
|
||||
error?: unknown;
|
||||
};
|
||||
|
||||
function RolloutTableStoryWrapper({
|
||||
maxWidth,
|
||||
rollouts = sampleRollouts,
|
||||
isFetching = false,
|
||||
isError = false,
|
||||
error = null,
|
||||
}: WrapperProps) {
|
||||
const [searchTerm, setSearchTerm] = useState('');
|
||||
const [statusFilters, setStatusFilters] = useState<RolloutStatus[]>([]);
|
||||
const [modeFilters, setModeFilters] = useState<RolloutMode[]>([]);
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(5);
|
||||
const [sort, setSort] = useState<RolloutsSortState>({
|
||||
column: 'startTimestamp',
|
||||
direction: 'desc',
|
||||
});
|
||||
|
||||
const tableRecords = useMemo<RolloutTableRecord[]>(() => {
|
||||
if (!rollouts) {
|
||||
return [];
|
||||
}
|
||||
return rollouts.map((rolloutItem) => buildRolloutRecord(rolloutItem));
|
||||
}, [rollouts]);
|
||||
|
||||
const filteredRecords = useMemo(() => {
|
||||
const normalizedSearch = searchTerm.trim().toLowerCase();
|
||||
return tableRecords.filter((record) => {
|
||||
const matchesSearch = normalizedSearch.length === 0 || record.rolloutId.toLowerCase().includes(normalizedSearch);
|
||||
const matchesStatus = statusFilters.length === 0 || statusFilters.includes(record.status);
|
||||
const matchesMode = modeFilters.length === 0 || (record.mode !== null && modeFilters.includes(record.mode));
|
||||
return matchesSearch && matchesStatus && matchesMode;
|
||||
});
|
||||
}, [modeFilters, searchTerm, statusFilters, tableRecords]);
|
||||
|
||||
const sortedRecords = useMemo(() => {
|
||||
const sorted = filteredRecords.slice();
|
||||
if (!sorted.length) {
|
||||
return sorted;
|
||||
}
|
||||
const comparatorKey = sort.column as keyof RolloutTableRecord;
|
||||
if (!(comparatorKey in sorted[0])) {
|
||||
return sorted;
|
||||
}
|
||||
sorted.sort((a, b) => compareRecords(a, b, comparatorKey));
|
||||
if (sort.direction === 'desc') {
|
||||
sorted.reverse();
|
||||
}
|
||||
return sorted;
|
||||
}, [filteredRecords, sort]);
|
||||
|
||||
const totalRecordsValue = sortedRecords.length;
|
||||
|
||||
const pagedRecords = useMemo(() => {
|
||||
const startIndex = (page - 1) * recordsPerPage;
|
||||
const endIndex = startIndex + recordsPerPage;
|
||||
return sortedRecords.slice(startIndex, endIndex);
|
||||
}, [page, recordsPerPage, sortedRecords]);
|
||||
|
||||
const pagedRollouts = useMemo(() => pagedRecords.map((record) => record as Rollout), [pagedRecords]);
|
||||
|
||||
return (
|
||||
<Box mx='auto' style={{ maxWidth, width: '100%', padding: 16 }}>
|
||||
<Stack gap='md'>
|
||||
<Title order={2}>Rollouts</Title>
|
||||
<TextInput
|
||||
placeholder='Search by Rollout ID'
|
||||
value={searchTerm}
|
||||
onChange={(event) => setSearchTerm(event.currentTarget.value)}
|
||||
leftSection={<IconSearch size={16} />}
|
||||
data-testid='rollouts-search-input'
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
<RolloutTable
|
||||
rollouts={pagedRollouts}
|
||||
totalRecords={totalRecordsValue}
|
||||
isFetching={isFetching}
|
||||
isError={isError}
|
||||
error={error}
|
||||
searchTerm={searchTerm}
|
||||
statusFilters={statusFilters}
|
||||
modeFilters={modeFilters}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onStatusFilterChange={(values) => {
|
||||
setStatusFilters(values);
|
||||
setPage(1);
|
||||
}}
|
||||
onStatusFilterReset={() => {
|
||||
setStatusFilters([]);
|
||||
setPage(1);
|
||||
}}
|
||||
onModeFilterChange={(values) => {
|
||||
setModeFilters(values);
|
||||
setPage(1);
|
||||
}}
|
||||
onModeFilterReset={() => {
|
||||
setModeFilters([]);
|
||||
setPage(1);
|
||||
}}
|
||||
onSortStatusChange={(status) => {
|
||||
setSort({
|
||||
column: status.columnAccessor as string,
|
||||
direction: status.direction,
|
||||
});
|
||||
}}
|
||||
onPageChange={setPage}
|
||||
onRecordsPerPageChange={(value) => {
|
||||
setRecordsPerPage(value);
|
||||
setPage(1);
|
||||
}}
|
||||
onResetFilters={() => {
|
||||
setSearchTerm('');
|
||||
setStatusFilters([]);
|
||||
setModeFilters([]);
|
||||
setSort({ column: 'startTimestamp', direction: 'desc' });
|
||||
setPage(1);
|
||||
}}
|
||||
onRefetch={() => undefined}
|
||||
recordsPerPageOptions={[5, 10, 20]}
|
||||
/>
|
||||
</Stack>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
export const WideContainer: Story = {
|
||||
render: () => <RolloutTableStoryWrapper maxWidth={1280} />,
|
||||
};
|
||||
|
||||
export const MediumContainer: Story = {
|
||||
render: () => <RolloutTableStoryWrapper maxWidth={960} />,
|
||||
};
|
||||
|
||||
export const NarrowContainer: Story = {
|
||||
render: () => <RolloutTableStoryWrapper maxWidth={720} />,
|
||||
};
|
||||
|
||||
export const DrawerWidth: Story = {
|
||||
render: () => <RolloutTableStoryWrapper maxWidth={520} />,
|
||||
};
|
||||
|
||||
export const ErrorState: Story = {
|
||||
render: () => (
|
||||
<RolloutTableStoryWrapper maxWidth={600} rollouts={[]} isError error={new Error('Network unreachable')} />
|
||||
),
|
||||
};
|
||||
@@ -0,0 +1,469 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo } from 'react';
|
||||
import {
|
||||
IconAlertCircle,
|
||||
IconCheck,
|
||||
IconCopy,
|
||||
IconFileDescription,
|
||||
IconRefresh,
|
||||
IconRouteSquare,
|
||||
} from '@tabler/icons-react';
|
||||
import { DataTable, type DataTableColumn, type DataTableSortStatus } from 'mantine-datatable';
|
||||
import { ActionIcon, Badge, Box, Button, CopyButton, Group, Stack, Text, Tooltip } from '@mantine/core';
|
||||
import { useElementSize, useViewportSize } from '@mantine/hooks';
|
||||
import { getLayoutAwareWidth } from '@/layouts/helper';
|
||||
import type { Span } from '@/types';
|
||||
import { getErrorDescriptor } from '@/utils/error';
|
||||
import { formatDateTimeWithMilliseconds, formatDuration, toTimestamp } from '@/utils/format';
|
||||
import { createResponsiveColumns, type ColumnVisibilityConfig } from '@/utils/table';
|
||||
|
||||
const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
name: { minWidth: 12.5, priority: 0 },
|
||||
sequenceId: { fixedWidth: 6, priority: 1 },
|
||||
spanId: { fixedWidth: 14, priority: 1 },
|
||||
traceId: { fixedWidth: 24, priority: 3 },
|
||||
parentId: { fixedWidth: 12, priority: 2 },
|
||||
statusCode: { fixedWidth: 8, priority: 2 },
|
||||
attributeKeys: { minWidth: 12.5, priority: 2 },
|
||||
startTime: { fixedWidth: 15, priority: 1 },
|
||||
endTime: { fixedWidth: 15, priority: 1 },
|
||||
duration: { fixedWidth: 10, priority: 3 },
|
||||
actionsPlaceholder: { fixedWidth: 6, priority: 0 },
|
||||
};
|
||||
|
||||
const STATUS_COLORS: Record<string, string> = {
|
||||
UNSET: 'gray',
|
||||
OK: 'teal',
|
||||
ERROR: 'red',
|
||||
};
|
||||
|
||||
export type TracesTableRecord = Span & {
|
||||
statusCode: string;
|
||||
attributeKeys: string;
|
||||
duration: number;
|
||||
actionsPlaceholder?: null;
|
||||
};
|
||||
|
||||
export function buildTraceRecord(span: Span): TracesTableRecord {
|
||||
const statusCode = span.status.status_code;
|
||||
const attributeKeys = Object.keys(span.attributes ?? {}).join(', ') || '';
|
||||
const startTimestamp = toTimestamp(span.startTime);
|
||||
const endTimestamp = toTimestamp(span.endTime);
|
||||
const duration = endTimestamp && startTimestamp ? endTimestamp - startTimestamp : 0;
|
||||
|
||||
return {
|
||||
...span,
|
||||
statusCode,
|
||||
attributeKeys,
|
||||
duration,
|
||||
actionsPlaceholder: null,
|
||||
};
|
||||
}
|
||||
|
||||
type TracesColumnsOptions = {
|
||||
onShowRollout?: (record: TracesTableRecord) => void;
|
||||
onShowSpanDetail?: (record: TracesTableRecord) => void;
|
||||
onParentIdClick?: (parentId: string) => void;
|
||||
spanIds: Set<string>;
|
||||
};
|
||||
|
||||
function createTracesColumns({
|
||||
onShowRollout,
|
||||
onShowSpanDetail,
|
||||
onParentIdClick,
|
||||
spanIds,
|
||||
}: TracesColumnsOptions): DataTableColumn<TracesTableRecord>[] {
|
||||
return [
|
||||
{
|
||||
accessor: 'name',
|
||||
title: 'Name',
|
||||
sortable: true,
|
||||
render: ({ name }) => (
|
||||
<Text size='sm' fw={500}>
|
||||
{name}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'sequenceId',
|
||||
title: 'Seq.',
|
||||
sortable: true,
|
||||
render: ({ sequenceId }) => <Text size='sm'>{sequenceId}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'traceId',
|
||||
title: 'Trace ID',
|
||||
sortable: true,
|
||||
render: ({ traceId }) => (
|
||||
<Group gap={2}>
|
||||
<Text size='sm'>{traceId}</Text>
|
||||
<CopyButton value={traceId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy trace ID ${traceId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'spanId',
|
||||
title: 'Span ID',
|
||||
sortable: true,
|
||||
render: ({ spanId }) => (
|
||||
<Group gap={2}>
|
||||
<Text size='sm'>{spanId}</Text>
|
||||
<CopyButton value={spanId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy span ID ${spanId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'parentId',
|
||||
title: 'Parent ID',
|
||||
sortable: true,
|
||||
render: ({ parentId }) => {
|
||||
if (!parentId) {
|
||||
return (
|
||||
<Text size='sm' c='dimmed'>
|
||||
—
|
||||
</Text>
|
||||
);
|
||||
}
|
||||
|
||||
const parentExists = spanIds.has(parentId);
|
||||
const isInteractive = parentExists && typeof onParentIdClick === 'function';
|
||||
|
||||
return (
|
||||
<Group gap={2}>
|
||||
<Text
|
||||
size='sm'
|
||||
c={parentExists ? undefined : 'red'}
|
||||
style={{ cursor: isInteractive ? 'pointer' : undefined }}
|
||||
onClick={(event) => {
|
||||
if (isInteractive) {
|
||||
event.stopPropagation();
|
||||
onParentIdClick?.(parentId);
|
||||
}
|
||||
}}
|
||||
>
|
||||
{parentId.slice(0, 8)}
|
||||
</Text>
|
||||
{!parentExists && (
|
||||
<Tooltip label='Parent span not found in table' withArrow>
|
||||
<IconAlertCircle size={14} color='red' />
|
||||
</Tooltip>
|
||||
)}
|
||||
</Group>
|
||||
);
|
||||
},
|
||||
},
|
||||
|
||||
{
|
||||
accessor: 'statusCode',
|
||||
title: 'Status',
|
||||
sortable: true,
|
||||
render: ({ statusCode }) => (
|
||||
<Badge size='sm' variant='light' color={STATUS_COLORS[statusCode] ?? 'gray'}>
|
||||
{statusCode}
|
||||
</Badge>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'attributeKeys',
|
||||
title: 'Attribute Keys',
|
||||
render: ({ attributeKeys }) =>
|
||||
attributeKeys ? (
|
||||
<Text size='sm' lineClamp={1}>
|
||||
{/* TODO: dim "." and "," and other characters are just normal text */}
|
||||
{attributeKeys}
|
||||
</Text>
|
||||
) : (
|
||||
<Text size='sm' c='dimmed'>
|
||||
—
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'startTime',
|
||||
title: 'Start Time',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ startTime }) => <Text size='sm'>{formatDateTimeWithMilliseconds(toTimestamp(startTime))}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'endTime',
|
||||
title: 'End Time',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ endTime }) => <Text size='sm'>{formatDateTimeWithMilliseconds(toTimestamp(endTime))}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'duration',
|
||||
title: 'Duration',
|
||||
sortable: true,
|
||||
textAlign: 'left',
|
||||
render: ({ duration }) => <Text size='sm'>{formatDuration(duration)}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'actionsPlaceholder',
|
||||
title: 'Actions',
|
||||
render: (record) => (
|
||||
<Group gap={2}>
|
||||
<Tooltip label='Show rollout' withArrow disabled={!onShowRollout}>
|
||||
<ActionIcon
|
||||
aria-label='Show rollout'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onShowRollout?.(record);
|
||||
}}
|
||||
>
|
||||
<IconRouteSquare size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
<Tooltip label='Show span detail' withArrow disabled={!onShowSpanDetail}>
|
||||
<ActionIcon
|
||||
aria-label='Show span detail'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onShowSpanDetail?.(record);
|
||||
}}
|
||||
>
|
||||
<IconFileDescription size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
export type TracesTableProps = {
|
||||
spans: Span[] | undefined;
|
||||
totalRecords: number;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
error: unknown;
|
||||
selectionMessage?: string;
|
||||
searchTerm: string;
|
||||
sort: { column: string; direction: 'asc' | 'desc' };
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
onSortStatusChange: (status: DataTableSortStatus<TracesTableRecord>) => void;
|
||||
onPageChange: (page: number) => void;
|
||||
onRecordsPerPageChange: (value: number) => void;
|
||||
onResetFilters: () => void;
|
||||
onRefetch: () => void;
|
||||
onShowRollout?: (record: TracesTableRecord) => void;
|
||||
onShowSpanDetail?: (record: TracesTableRecord) => void;
|
||||
onParentIdClick?: (parentId: string) => void;
|
||||
recordsPerPageOptions?: number[];
|
||||
};
|
||||
|
||||
export function TracesTable({
|
||||
spans,
|
||||
totalRecords,
|
||||
isFetching,
|
||||
isError,
|
||||
error,
|
||||
selectionMessage,
|
||||
searchTerm,
|
||||
sort,
|
||||
page,
|
||||
recordsPerPage,
|
||||
onSortStatusChange,
|
||||
onPageChange,
|
||||
onRecordsPerPageChange,
|
||||
onResetFilters,
|
||||
onRefetch,
|
||||
onShowRollout,
|
||||
onShowSpanDetail,
|
||||
onParentIdClick,
|
||||
recordsPerPageOptions = DEFAULT_RECORDS_PER_PAGE_OPTIONS,
|
||||
}: TracesTableProps) {
|
||||
const { ref: tableContainerRef, width: containerWidth } = useElementSize();
|
||||
const { width: viewportWidth } = useViewportSize();
|
||||
|
||||
const traceRecords = useMemo<TracesTableRecord[]>(() => {
|
||||
if (!spans) {
|
||||
return [];
|
||||
}
|
||||
return spans.map((span) => buildTraceRecord(span));
|
||||
}, [spans]);
|
||||
|
||||
const spanIds = useMemo(() => {
|
||||
return new Set(traceRecords.map((record) => record.spanId));
|
||||
}, [traceRecords]);
|
||||
|
||||
const columns = useMemo(
|
||||
() =>
|
||||
createTracesColumns({
|
||||
onShowRollout,
|
||||
onShowSpanDetail,
|
||||
onParentIdClick,
|
||||
spanIds,
|
||||
}),
|
||||
[onShowRollout, onShowSpanDetail, onParentIdClick, spanIds],
|
||||
);
|
||||
|
||||
const layoutAwareContainerWidth = useMemo(
|
||||
() => getLayoutAwareWidth(containerWidth, viewportWidth),
|
||||
[containerWidth, viewportWidth],
|
||||
);
|
||||
|
||||
const responsiveColumns = useMemo(
|
||||
() => createResponsiveColumns(columns, layoutAwareContainerWidth, COLUMN_VISIBILITY),
|
||||
[columns, layoutAwareContainerWidth],
|
||||
);
|
||||
|
||||
const totalPages = useMemo(
|
||||
() => Math.max(1, Math.ceil(Math.max(0, totalRecords) / Math.max(1, recordsPerPage))),
|
||||
[recordsPerPage, totalRecords],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (page > totalPages) {
|
||||
onPageChange(totalPages);
|
||||
}
|
||||
}, [onPageChange, page, totalPages]);
|
||||
|
||||
const hasActiveFilters = searchTerm.trim().length > 0;
|
||||
|
||||
const sortStatus: DataTableSortStatus<TracesTableRecord> = {
|
||||
columnAccessor: sort.column,
|
||||
direction: sort.direction,
|
||||
};
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<TracesTableRecord>) => {
|
||||
onSortStatusChange(status);
|
||||
},
|
||||
[onSortStatusChange],
|
||||
);
|
||||
|
||||
const errorDescriptor = isError ? getErrorDescriptor(error) : null;
|
||||
const errorMessage = isError
|
||||
? `Traces are temporarily unavailable${errorDescriptor ? ` (${errorDescriptor})` : ''}.`
|
||||
: 'Traces are temporarily unavailable.';
|
||||
|
||||
const selectionEmptyState = selectionMessage ? (
|
||||
<Stack gap='sm' align='center' py='xl'>
|
||||
<Text fw={600} size='sm'>
|
||||
{selectionMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Choose a rollout and attempt from the controls above to load trace results.
|
||||
</Text>
|
||||
</Stack>
|
||||
) : null;
|
||||
|
||||
const fallbackEmptyState = (
|
||||
<Stack gap='sm' align='center' py='lg'>
|
||||
{isError ? (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
{errorMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Use the retry button to try again, or adjust the filters to broaden the results.
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' color='gray' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Retry
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
No traces found
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
{hasActiveFilters
|
||||
? 'Try adjusting the search to see more results.'
|
||||
: 'Try refreshing to fetch the latest traces.'}
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Refresh
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
|
||||
const emptyState = selectionEmptyState ?? fallbackEmptyState;
|
||||
|
||||
return (
|
||||
<Box ref={tableContainerRef}>
|
||||
<DataTable<TracesTableRecord>
|
||||
classNames={{ root: 'traces-table' }}
|
||||
withTableBorder
|
||||
withColumnBorders
|
||||
highlightOnHover
|
||||
verticalAlign='center'
|
||||
minHeight={traceRecords.length === 0 ? 500 : undefined}
|
||||
idAccessor='spanId'
|
||||
records={traceRecords}
|
||||
columns={responsiveColumns}
|
||||
totalRecords={totalRecords}
|
||||
recordsPerPage={recordsPerPage}
|
||||
page={page}
|
||||
onPageChange={onPageChange}
|
||||
onRecordsPerPageChange={onRecordsPerPageChange}
|
||||
recordsPerPageOptions={recordsPerPageOptions}
|
||||
sortStatus={sortStatus}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
emptyState={traceRecords.length === 0 ? emptyState : undefined}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,367 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo, useState } from 'react';
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import { Box, Stack, TextInput, Title } from '@mantine/core';
|
||||
import type { Span } from '@/types';
|
||||
import { compareRecords } from '@/utils/table';
|
||||
import { buildTraceRecord, TracesTable, type TracesTableRecord } from './TracesTable.component';
|
||||
|
||||
const meta: Meta<typeof TracesTable> = {
|
||||
title: 'Components/TracesTable',
|
||||
component: TracesTable,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof TracesTable>;
|
||||
|
||||
const now = Math.floor(1762775145209 / 1000);
|
||||
|
||||
const sampleSpans: Span[] = [
|
||||
{
|
||||
rolloutId: 'ro-trace-001',
|
||||
attemptId: 'at-trace-001',
|
||||
sequenceId: 1,
|
||||
traceId: 'trace-abc123def456',
|
||||
spanId: 'span-root-001',
|
||||
parentId: null,
|
||||
name: 'main_task',
|
||||
status: { status_code: 'OK', description: null },
|
||||
attributes: {
|
||||
'task.type': 'generation',
|
||||
'task.priority': 'high',
|
||||
'user.id': 'user-123',
|
||||
},
|
||||
startTime: now - 100,
|
||||
endTime: now - 10,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-001',
|
||||
attemptId: 'at-trace-001',
|
||||
sequenceId: 1,
|
||||
traceId: 'trace-abc123def456',
|
||||
spanId: 'span-child-001',
|
||||
parentId: 'span-root-001',
|
||||
name: 'llm_call',
|
||||
status: { status_code: 'OK', description: null },
|
||||
attributes: {
|
||||
'llm.model': 'gpt-4',
|
||||
'llm.temperature': 0.7,
|
||||
'llm.max_tokens': 2048,
|
||||
},
|
||||
startTime: now - 90,
|
||||
endTime: now - 50,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-001',
|
||||
attemptId: 'at-trace-001',
|
||||
sequenceId: 1,
|
||||
traceId: 'trace-abc123def456',
|
||||
spanId: 'span-child-002',
|
||||
parentId: 'span-root-001',
|
||||
name: 'database_query',
|
||||
status: { status_code: 'OK', description: null },
|
||||
attributes: {
|
||||
'db.system': 'postgresql',
|
||||
'db.operation': 'SELECT',
|
||||
'db.table': 'users',
|
||||
},
|
||||
startTime: now - 80,
|
||||
endTime: now - 70,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-002',
|
||||
attemptId: 'at-trace-002',
|
||||
sequenceId: 1,
|
||||
traceId: 'trace-xyz789ghi012',
|
||||
spanId: 'span-error-001',
|
||||
parentId: 'span-missing-parent',
|
||||
name: 'failed_operation',
|
||||
status: { status_code: 'ERROR', description: 'Connection timeout' },
|
||||
attributes: {
|
||||
'error.type': 'TimeoutError',
|
||||
'error.message': 'Connection timed out after 30s',
|
||||
},
|
||||
startTime: now - 150,
|
||||
endTime: now - 120,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-003',
|
||||
attemptId: 'at-trace-003',
|
||||
sequenceId: 1,
|
||||
traceId: 'trace-unset123',
|
||||
spanId: 'span-unset-001',
|
||||
parentId: null,
|
||||
name: 'pending_task',
|
||||
status: { status_code: 'UNSET', description: null },
|
||||
attributes: {},
|
||||
startTime: now - 30,
|
||||
endTime: now - 5,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-004',
|
||||
attemptId: 'at-trace-004',
|
||||
sequenceId: 2,
|
||||
traceId: 'trace-nested456',
|
||||
spanId: 'span-parent-001',
|
||||
parentId: null,
|
||||
name: 'workflow_execution',
|
||||
status: { status_code: 'OK', description: null },
|
||||
attributes: {
|
||||
'workflow.name': 'data_processing',
|
||||
'workflow.version': '2.1.0',
|
||||
},
|
||||
startTime: now - 200,
|
||||
endTime: now - 50,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-004',
|
||||
attemptId: 'at-trace-004',
|
||||
sequenceId: 2,
|
||||
traceId: 'trace-nested456',
|
||||
spanId: 'span-child-nested-001',
|
||||
parentId: 'span-parent-001',
|
||||
name: 'step_1_validation',
|
||||
status: { status_code: 'OK', description: null },
|
||||
attributes: {
|
||||
'step.name': 'validation',
|
||||
'step.index': 1,
|
||||
},
|
||||
startTime: now - 195,
|
||||
endTime: now - 180,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
{
|
||||
rolloutId: 'ro-trace-004',
|
||||
attemptId: 'at-trace-004',
|
||||
sequenceId: 2,
|
||||
traceId: 'trace-nested456',
|
||||
spanId: 'span-child-nested-002',
|
||||
parentId: 'span-parent-001',
|
||||
name: 'step_2_processing',
|
||||
status: { status_code: 'ERROR', description: 'Validation failed' },
|
||||
attributes: {
|
||||
'step.name': 'processing',
|
||||
'step.index': 2,
|
||||
'error.type': 'ValidationError',
|
||||
},
|
||||
startTime: now - 175,
|
||||
endTime: now - 160,
|
||||
events: [],
|
||||
links: [],
|
||||
context: {},
|
||||
parent: null,
|
||||
resource: {},
|
||||
},
|
||||
];
|
||||
|
||||
type WrapperProps = {
|
||||
maxWidth: number;
|
||||
spans?: Span[] | undefined;
|
||||
isFetching?: boolean;
|
||||
isError?: boolean;
|
||||
error?: unknown;
|
||||
};
|
||||
|
||||
function TracesTableStoryWrapper({
|
||||
maxWidth,
|
||||
spans = sampleSpans,
|
||||
isFetching = false,
|
||||
isError = false,
|
||||
error = null,
|
||||
}: WrapperProps) {
|
||||
const [searchTerm, setSearchTerm] = useState('');
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(10);
|
||||
const [sort, setSort] = useState<{ column: string; direction: 'asc' | 'desc' }>({
|
||||
column: 'startTime',
|
||||
direction: 'desc',
|
||||
});
|
||||
|
||||
const tableRecords = useMemo<TracesTableRecord[]>(() => {
|
||||
if (!spans) {
|
||||
return [];
|
||||
}
|
||||
return spans.map((span) => buildTraceRecord(span));
|
||||
}, [spans]);
|
||||
|
||||
const filteredRecords = useMemo(() => {
|
||||
const normalizedSearch = searchTerm.trim().toLowerCase();
|
||||
if (normalizedSearch.length === 0) {
|
||||
return tableRecords;
|
||||
}
|
||||
return tableRecords.filter(
|
||||
(record) =>
|
||||
record.traceId.toLowerCase().includes(normalizedSearch) ||
|
||||
record.spanId.toLowerCase().includes(normalizedSearch) ||
|
||||
record.name.toLowerCase().includes(normalizedSearch),
|
||||
);
|
||||
}, [searchTerm, tableRecords]);
|
||||
|
||||
const sortedRecords = useMemo(() => {
|
||||
const sorted = filteredRecords.slice();
|
||||
if (!sorted.length) {
|
||||
return sorted;
|
||||
}
|
||||
const comparatorKey = sort.column as keyof TracesTableRecord;
|
||||
if (!(comparatorKey in sorted[0])) {
|
||||
return sorted;
|
||||
}
|
||||
sorted.sort((a, b) => compareRecords(a, b, comparatorKey));
|
||||
if (sort.direction === 'desc') {
|
||||
sorted.reverse();
|
||||
}
|
||||
return sorted;
|
||||
}, [filteredRecords, sort]);
|
||||
|
||||
const totalRecordsValue = sortedRecords.length;
|
||||
|
||||
const pagedRecords = useMemo(() => {
|
||||
const startIndex = (page - 1) * recordsPerPage;
|
||||
const endIndex = startIndex + recordsPerPage;
|
||||
return sortedRecords.slice(startIndex, endIndex);
|
||||
}, [page, recordsPerPage, sortedRecords]);
|
||||
|
||||
const pagedSpans = useMemo(() => pagedRecords.map((record) => record as Span), [pagedRecords]);
|
||||
|
||||
const handleShowRollout = (record: any) => {
|
||||
console.log('Show rollout for:', record.rolloutId);
|
||||
};
|
||||
|
||||
const handleShowSpanDetail = (record: any) => {
|
||||
console.log('Show span detail for:', record.spanId, record);
|
||||
};
|
||||
|
||||
const handleParentIdClick = (parentId: string) => {
|
||||
console.log('Navigate to parent span:', parentId);
|
||||
setSearchTerm(parentId);
|
||||
};
|
||||
|
||||
return (
|
||||
<Box mx='auto' style={{ maxWidth, width: '100%', padding: 16 }}>
|
||||
<Stack gap='md'>
|
||||
<Title order={2}>Traces</Title>
|
||||
<TextInput
|
||||
placeholder='Search by Trace ID, Span ID, or Name'
|
||||
value={searchTerm}
|
||||
onChange={(event) => setSearchTerm(event.currentTarget.value)}
|
||||
leftSection={<IconSearch size={16} />}
|
||||
data-testid='traces-search-input'
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
<TracesTable
|
||||
spans={pagedSpans}
|
||||
totalRecords={totalRecordsValue}
|
||||
isFetching={isFetching}
|
||||
isError={isError}
|
||||
error={error}
|
||||
searchTerm={searchTerm}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={(status) => {
|
||||
setSort({
|
||||
column: status.columnAccessor as string,
|
||||
direction: status.direction,
|
||||
});
|
||||
}}
|
||||
onPageChange={setPage}
|
||||
onRecordsPerPageChange={(value) => {
|
||||
setRecordsPerPage(value);
|
||||
setPage(1);
|
||||
}}
|
||||
onResetFilters={() => {
|
||||
setSearchTerm('');
|
||||
setSort({ column: 'startTime', direction: 'desc' });
|
||||
setPage(1);
|
||||
}}
|
||||
onRefetch={() => undefined}
|
||||
onShowRollout={handleShowRollout}
|
||||
onShowSpanDetail={handleShowSpanDetail}
|
||||
onParentIdClick={handleParentIdClick}
|
||||
recordsPerPageOptions={[10, 20, 50]}
|
||||
/>
|
||||
</Stack>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
|
||||
export const WideContainer: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={1400} />,
|
||||
};
|
||||
|
||||
export const MediumContainer: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={960} />,
|
||||
};
|
||||
|
||||
export const NarrowContainer: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={720} />,
|
||||
};
|
||||
|
||||
export const DrawerWidth: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={520} />,
|
||||
};
|
||||
|
||||
export const ErrorState: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={960} spans={[]} isError error={new Error('Network unreachable')} />,
|
||||
};
|
||||
|
||||
export const LoadingState: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={960} spans={[]} isFetching />,
|
||||
};
|
||||
|
||||
export const EmptyState: Story = {
|
||||
render: () => <TracesTableStoryWrapper maxWidth={960} spans={[]} />,
|
||||
};
|
||||
|
||||
export const WithMissingParent: Story = {
|
||||
render: () => (
|
||||
<TracesTableStoryWrapper maxWidth={1200} spans={sampleSpans.filter((s) => s.spanId === 'span-error-001')} />
|
||||
),
|
||||
};
|
||||
|
||||
export const NestedSpans: Story = {
|
||||
render: () => (
|
||||
<TracesTableStoryWrapper maxWidth={1200} spans={sampleSpans.filter((s) => s.traceId === 'trace-nested456')} />
|
||||
),
|
||||
};
|
||||
@@ -0,0 +1,362 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useCallback, useEffect, useMemo } from 'react';
|
||||
import { IconCheck, IconCopy, IconInfoCircle, IconRefresh } from '@tabler/icons-react';
|
||||
import { DataTable, type DataTableColumn, type DataTableSortStatus } from 'mantine-datatable';
|
||||
import { ActionIcon, Badge, Box, Button, CopyButton, Group, Stack, Text, Tooltip } from '@mantine/core';
|
||||
import { useElementSize, useViewportSize } from '@mantine/hooks';
|
||||
import { getLayoutAwareWidth } from '@/layouts/helper';
|
||||
import type { Worker } from '@/types';
|
||||
import { getErrorDescriptor } from '@/utils/error';
|
||||
import { formatDateTime, formatRelativeTime, formatStatusLabel } from '@/utils/format';
|
||||
import { createResponsiveColumns, type ColumnVisibilityConfig } from '@/utils/table';
|
||||
|
||||
const DEFAULT_RECORDS_PER_PAGE_OPTIONS = [50, 100, 200, 500];
|
||||
|
||||
const COLUMN_VISIBILITY: Record<string, ColumnVisibilityConfig> = {
|
||||
workerId: { fixedWidth: 12, priority: 0 },
|
||||
status: { fixedWidth: 6, priority: 1 },
|
||||
currentRolloutId: { fixedWidth: 14, priority: 3 },
|
||||
currentAttemptId: { fixedWidth: 14, priority: 3 },
|
||||
lastHeartbeatTime: { fixedWidth: 10, priority: 2 },
|
||||
lastBusyTime: { fixedWidth: 10, priority: 3 },
|
||||
lastIdleTime: { fixedWidth: 10, priority: 3 },
|
||||
lastDequeueTime: { fixedWidth: 10, priority: 1 },
|
||||
actions: { fixedWidth: 5, priority: 0 },
|
||||
};
|
||||
|
||||
export type WorkersTableRecord = Worker & {
|
||||
timestamps: Record<
|
||||
'lastHeartbeatTime' | 'lastBusyTime' | 'lastIdleTime' | 'lastDequeueTime',
|
||||
{ absolute: string; relative: string }
|
||||
>;
|
||||
};
|
||||
|
||||
const buildTimestampMeta = (value: Worker['lastHeartbeatTime']) => ({
|
||||
absolute: formatDateTime(value),
|
||||
relative: formatRelativeTime(value),
|
||||
});
|
||||
|
||||
function buildWorkerRecord(worker: Worker): WorkersTableRecord {
|
||||
return {
|
||||
...worker,
|
||||
timestamps: {
|
||||
lastHeartbeatTime: buildTimestampMeta(worker.lastHeartbeatTime),
|
||||
lastBusyTime: buildTimestampMeta(worker.lastBusyTime),
|
||||
lastIdleTime: buildTimestampMeta(worker.lastIdleTime),
|
||||
lastDequeueTime: buildTimestampMeta(worker.lastDequeueTime),
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
type WorkersColumnsOptions = {
|
||||
onShowDetails: (worker: Worker) => void;
|
||||
};
|
||||
|
||||
const STATUS_COLORS: Record<Worker['status'], string> = {
|
||||
busy: 'orange',
|
||||
idle: 'teal',
|
||||
unknown: 'gray',
|
||||
};
|
||||
|
||||
function createWorkersColumns({ onShowDetails }: WorkersColumnsOptions): DataTableColumn<WorkersTableRecord>[] {
|
||||
return [
|
||||
{
|
||||
accessor: 'workerId',
|
||||
title: 'Runner ID',
|
||||
sortable: true,
|
||||
render: ({ workerId }) => (
|
||||
<Group gap={2} wrap='nowrap'>
|
||||
<Text fw={500} size='sm'>
|
||||
{workerId}
|
||||
</Text>
|
||||
<CopyButton value={workerId}>
|
||||
{({ copied, copy }) => (
|
||||
<Tooltip label={copied ? 'Copied' : 'Copy'} withArrow>
|
||||
<ActionIcon
|
||||
aria-label={`Copy worker ID ${workerId}`}
|
||||
variant='subtle'
|
||||
color={copied ? 'teal' : 'gray'}
|
||||
size='sm'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
copy();
|
||||
}}
|
||||
>
|
||||
{copied ? <IconCheck size={14} /> : <IconCopy size={14} />}
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CopyButton>
|
||||
</Group>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'status',
|
||||
title: 'Status',
|
||||
sortable: true,
|
||||
render: ({ status }) => {
|
||||
const color = STATUS_COLORS[status] ?? 'gray';
|
||||
return (
|
||||
<Badge size='sm' variant='light' color={color} radius='sm'>
|
||||
{formatStatusLabel(status)}
|
||||
</Badge>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
accessor: 'currentRolloutId',
|
||||
title: 'Current Rollout',
|
||||
sortable: true,
|
||||
render: ({ currentRolloutId }) => <Text size='sm'>{currentRolloutId ?? '—'}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'currentAttemptId',
|
||||
title: 'Current Attempt',
|
||||
sortable: true,
|
||||
render: ({ currentAttemptId }) => <Text size='sm'>{currentAttemptId ?? '—'}</Text>,
|
||||
},
|
||||
{
|
||||
accessor: 'lastHeartbeatTime',
|
||||
title: 'Heartbeat',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastHeartbeatTime.relative}</Text>
|
||||
{timestamps.lastHeartbeatTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastHeartbeatTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastBusyTime',
|
||||
title: 'Last Busy',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastBusyTime.relative}</Text>
|
||||
{timestamps.lastBusyTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastBusyTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastIdleTime',
|
||||
title: 'Last Idle',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastIdleTime.relative}</Text>
|
||||
{timestamps.lastIdleTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastIdleTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'lastDequeueTime',
|
||||
title: 'Last Dequeue',
|
||||
sortable: true,
|
||||
render: ({ timestamps }) => (
|
||||
<Stack gap={0} justify='center'>
|
||||
<Text size='sm'>{timestamps.lastDequeueTime.relative}</Text>
|
||||
{timestamps.lastDequeueTime.absolute !== '—' && (
|
||||
<Text size='xs' c='dimmed'>
|
||||
{timestamps.lastDequeueTime.absolute}
|
||||
</Text>
|
||||
)}
|
||||
</Stack>
|
||||
),
|
||||
},
|
||||
{
|
||||
accessor: 'actions',
|
||||
title: 'Actions',
|
||||
textAlign: 'left',
|
||||
render: (record) => (
|
||||
<Tooltip label='Show runner detail' withArrow disabled={!onShowDetails}>
|
||||
<ActionIcon
|
||||
aria-label='Show runner detail'
|
||||
variant='subtle'
|
||||
color='gray'
|
||||
onClick={(event) => {
|
||||
event.stopPropagation();
|
||||
onShowDetails(record);
|
||||
}}
|
||||
>
|
||||
<IconInfoCircle size={16} />
|
||||
</ActionIcon>
|
||||
</Tooltip>
|
||||
),
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
export type WorkersTableProps = {
|
||||
workers: Worker[] | undefined;
|
||||
totalRecords: number;
|
||||
isFetching: boolean;
|
||||
isError: boolean;
|
||||
error: unknown;
|
||||
searchTerm: string;
|
||||
sort: { column: string; direction: 'asc' | 'desc' };
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
onSortStatusChange: (status: DataTableSortStatus<WorkersTableRecord>) => void;
|
||||
onPageChange: (page: number) => void;
|
||||
onRecordsPerPageChange: (value: number) => void;
|
||||
onResetFilters: () => void;
|
||||
onRefetch: () => void;
|
||||
onShowDetails: (worker: Worker) => void;
|
||||
recordsPerPageOptions?: number[];
|
||||
};
|
||||
|
||||
export function WorkersTable({
|
||||
workers,
|
||||
totalRecords,
|
||||
isFetching,
|
||||
isError,
|
||||
error,
|
||||
searchTerm,
|
||||
sort,
|
||||
page,
|
||||
recordsPerPage,
|
||||
onSortStatusChange,
|
||||
onPageChange,
|
||||
onRecordsPerPageChange,
|
||||
onResetFilters,
|
||||
onRefetch,
|
||||
onShowDetails,
|
||||
recordsPerPageOptions = DEFAULT_RECORDS_PER_PAGE_OPTIONS,
|
||||
}: WorkersTableProps) {
|
||||
const { ref: tableContainerRef, width: containerWidth } = useElementSize();
|
||||
const { width: viewportWidth } = useViewportSize();
|
||||
|
||||
const layoutAwareContainerWidth = useMemo(
|
||||
() => getLayoutAwareWidth(containerWidth, viewportWidth),
|
||||
[containerWidth, viewportWidth],
|
||||
);
|
||||
|
||||
const workerRecords = useMemo<WorkersTableRecord[]>(() => {
|
||||
if (!workers) {
|
||||
return [];
|
||||
}
|
||||
return workers.map((worker) => buildWorkerRecord(worker));
|
||||
}, [workers]);
|
||||
|
||||
const columns = useMemo(() => createWorkersColumns({ onShowDetails }), [onShowDetails]);
|
||||
const responsiveColumns = useMemo(
|
||||
() => createResponsiveColumns(columns, layoutAwareContainerWidth, COLUMN_VISIBILITY),
|
||||
[columns, layoutAwareContainerWidth],
|
||||
);
|
||||
|
||||
const totalPages = useMemo(
|
||||
() => Math.max(1, Math.ceil(Math.max(0, totalRecords) / Math.max(1, recordsPerPage))),
|
||||
[recordsPerPage, totalRecords],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (page > totalPages) {
|
||||
onPageChange(totalPages);
|
||||
}
|
||||
}, [onPageChange, page, totalPages]);
|
||||
|
||||
const hasActiveFilters = searchTerm.trim().length > 0;
|
||||
|
||||
const sortStatus: DataTableSortStatus<WorkersTableRecord> = {
|
||||
columnAccessor: sort.column,
|
||||
direction: sort.direction,
|
||||
};
|
||||
|
||||
const handleSortStatusChange = useCallback(
|
||||
(status: DataTableSortStatus<WorkersTableRecord>) => {
|
||||
onSortStatusChange(status);
|
||||
},
|
||||
[onSortStatusChange],
|
||||
);
|
||||
|
||||
const errorDescriptor = isError ? getErrorDescriptor(error) : null;
|
||||
const errorMessage = isError
|
||||
? `Workers are temporarily unavailable${errorDescriptor ? ` (${errorDescriptor})` : ''}.`
|
||||
: 'Workers are temporarily unavailable.';
|
||||
|
||||
const emptyState = (
|
||||
<Stack gap='sm' align='center' py='lg'>
|
||||
{isError ? (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
{errorMessage}
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
Use the retry button to try again, or adjust the search to broaden the results.
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' color='gray' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Retry
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
<Text fw={600} size='sm'>
|
||||
No workers found
|
||||
</Text>
|
||||
<Text size='sm' c='dimmed' ta='center'>
|
||||
{hasActiveFilters
|
||||
? 'Try adjusting the search to see more results.'
|
||||
: 'Try refreshing to fetch the latest worker status.'}
|
||||
</Text>
|
||||
<Group gap='xs'>
|
||||
<Button size='xs' variant='light' leftSection={<IconRefresh size={14} />} onClick={onRefetch}>
|
||||
Refresh
|
||||
</Button>
|
||||
{hasActiveFilters ? (
|
||||
<Button size='xs' variant='subtle' onClick={onResetFilters}>
|
||||
Clear filters
|
||||
</Button>
|
||||
) : null}
|
||||
</Group>
|
||||
</>
|
||||
)}
|
||||
</Stack>
|
||||
);
|
||||
|
||||
return (
|
||||
<Box ref={tableContainerRef}>
|
||||
<DataTable<WorkersTableRecord>
|
||||
classNames={{ root: 'workers-table' }}
|
||||
withTableBorder
|
||||
withColumnBorders
|
||||
highlightOnHover
|
||||
verticalAlign='center'
|
||||
minHeight={workerRecords.length === 0 ? 400 : undefined}
|
||||
idAccessor='workerId'
|
||||
records={workerRecords}
|
||||
columns={responsiveColumns}
|
||||
totalRecords={totalRecords}
|
||||
recordsPerPage={recordsPerPage}
|
||||
page={page}
|
||||
onPageChange={onPageChange}
|
||||
onRecordsPerPageChange={onRecordsPerPageChange}
|
||||
recordsPerPageOptions={recordsPerPageOptions}
|
||||
sortStatus={sortStatus}
|
||||
onSortStatusChange={handleSortStatusChange}
|
||||
fetching={isFetching}
|
||||
loaderSize='sm'
|
||||
emptyState={workerRecords.length === 0 ? emptyState : undefined}
|
||||
/>
|
||||
</Box>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { useMemo, useState } from 'react';
|
||||
import type { Meta, StoryObj } from '@storybook/react';
|
||||
import { IconSearch } from '@tabler/icons-react';
|
||||
import { Box, Stack, TextInput, Title } from '@mantine/core';
|
||||
import type { Worker } from '@/types';
|
||||
import { WorkersTable } from './WorkersTable.component';
|
||||
|
||||
const meta: Meta<typeof WorkersTable> = {
|
||||
title: 'Components/WorkersTable',
|
||||
component: WorkersTable,
|
||||
parameters: {
|
||||
layout: 'fullscreen',
|
||||
},
|
||||
};
|
||||
|
||||
export default meta;
|
||||
|
||||
type Story = StoryObj<typeof WorkersTable>;
|
||||
|
||||
const now = Math.floor(Date.now() / 1000);
|
||||
|
||||
const sampleWorkers: Worker[] = [
|
||||
{
|
||||
workerId: 'worker-east',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 2, gpuUtilization: 0.82 },
|
||||
lastHeartbeatTime: now - 20,
|
||||
lastDequeueTime: now - 60,
|
||||
lastBusyTime: now - 120,
|
||||
lastIdleTime: now - 600,
|
||||
currentRolloutId: 'ro-story-001',
|
||||
currentAttemptId: 'at-story-010',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-west',
|
||||
status: 'busy',
|
||||
heartbeatStats: { queueDepth: 1 },
|
||||
lastHeartbeatTime: now - 45,
|
||||
lastDequeueTime: now - 300,
|
||||
lastBusyTime: now - 200,
|
||||
lastIdleTime: now - 4800,
|
||||
currentRolloutId: 'ro-story-003',
|
||||
currentAttemptId: 'at-story-033',
|
||||
},
|
||||
{
|
||||
workerId: 'worker-north',
|
||||
status: 'idle',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 90,
|
||||
lastDequeueTime: now - 3600,
|
||||
lastBusyTime: now - 5400,
|
||||
lastIdleTime: now - 5400,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-south',
|
||||
status: 'idle',
|
||||
heartbeatStats: null,
|
||||
lastHeartbeatTime: now - 900,
|
||||
lastDequeueTime: now - 7200,
|
||||
lastBusyTime: now - 8600,
|
||||
lastIdleTime: now - 8600,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
{
|
||||
workerId: 'worker-standby',
|
||||
status: 'unknown',
|
||||
heartbeatStats: { queueDepth: 0 },
|
||||
lastHeartbeatTime: now - 15,
|
||||
lastDequeueTime: now - 4000,
|
||||
lastBusyTime: null,
|
||||
lastIdleTime: null,
|
||||
currentRolloutId: null,
|
||||
currentAttemptId: null,
|
||||
},
|
||||
];
|
||||
|
||||
type WorkersTableStoryWrapperProps = {
|
||||
maxWidth: number;
|
||||
initialSort?: { column: string; direction: 'asc' | 'desc' };
|
||||
};
|
||||
|
||||
function WorkersTableStoryWrapper({ maxWidth, initialSort }: WorkersTableStoryWrapperProps) {
|
||||
const [searchTerm, setSearchTerm] = useState('');
|
||||
const [page, setPage] = useState(1);
|
||||
const [recordsPerPage, setRecordsPerPage] = useState(5);
|
||||
const [sort, setSort] = useState<{ column: string; direction: 'asc' | 'desc' }>(
|
||||
() => initialSort ?? { column: 'lastHeartbeatTime', direction: 'desc' },
|
||||
);
|
||||
|
||||
const filteredWorkers = useMemo(() => {
|
||||
const normalized = searchTerm.trim().toLowerCase();
|
||||
if (!normalized) {
|
||||
return sampleWorkers;
|
||||
}
|
||||
return sampleWorkers.filter((worker) => worker.workerId.toLowerCase().includes(normalized));
|
||||
}, [searchTerm]);
|
||||
|
||||
return (
|
||||
<Stack gap='md' p='lg'>
|
||||
<Title order={2}>Workers ({maxWidth}px max width)</Title>
|
||||
<TextInput
|
||||
placeholder='Search'
|
||||
leftSection={<IconSearch size={16} />}
|
||||
value={searchTerm}
|
||||
onChange={(event) => setSearchTerm(event.currentTarget.value)}
|
||||
w='100%'
|
||||
style={{ maxWidth: 360 }}
|
||||
/>
|
||||
<Box style={{ maxWidth }}>
|
||||
<WorkersTable
|
||||
workers={filteredWorkers}
|
||||
totalRecords={filteredWorkers.length}
|
||||
isFetching={false}
|
||||
isError={false}
|
||||
error={null}
|
||||
searchTerm={searchTerm}
|
||||
sort={sort}
|
||||
page={page}
|
||||
recordsPerPage={recordsPerPage}
|
||||
onSortStatusChange={(status) => {
|
||||
return setSort({ column: status.columnAccessor as string, direction: status.direction });
|
||||
}}
|
||||
onPageChange={setPage}
|
||||
onRecordsPerPageChange={setRecordsPerPage}
|
||||
onResetFilters={() => {
|
||||
setSearchTerm('');
|
||||
setPage(1);
|
||||
}}
|
||||
onRefetch={() => {}}
|
||||
onShowDetails={() => {}}
|
||||
/>
|
||||
</Box>
|
||||
</Stack>
|
||||
);
|
||||
}
|
||||
|
||||
export const Wide: Story = {
|
||||
render: () => <WorkersTableStoryWrapper maxWidth={1600} />,
|
||||
};
|
||||
|
||||
export const Narrow: Story = {
|
||||
render: () => <WorkersTableStoryWrapper maxWidth={780} />,
|
||||
};
|
||||
|
||||
export const SortedByCurrentRollout: Story = {
|
||||
render: () => (
|
||||
<WorkersTableStoryWrapper maxWidth={1200} initialSort={{ column: 'currentRolloutId', direction: 'asc' }} />
|
||||
),
|
||||
};
|
||||
@@ -0,0 +1,436 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { alpha, CSSVariablesResolver } from '@mantine/core';
|
||||
|
||||
export const shadcnCssVariableResolver: CSSVariablesResolver = () => ({
|
||||
variables: {
|
||||
// variables that do not depend on color scheme
|
||||
'--mantine-heading-font-weight': '600',
|
||||
'--mantine-primary-color-filled-hover': alpha('var(--mantine-primary-color-filled)', 0.9),
|
||||
'--mantine-primary-color-light': 'var(--mantine-color-zinc-light)',
|
||||
'--mantine-primary-color-light-hover': 'var(--mantine-color-zinc-light-hover)',
|
||||
'--mantine-primary-color-light-color': 'var(--mantine-color-zinc-light-color)',
|
||||
},
|
||||
light: {
|
||||
// all variables that depend on light color scheme
|
||||
'--mantine-primary-color-contrast': 'var(--mantine-color-zinc-0)', // used as primary color contrast
|
||||
'--mantine-color-text': 'var(--mantine-color-secondary-9)', // used as text color
|
||||
'--mantine-color-body': 'var(--mantine-color-white)', // used as body color
|
||||
'--mantine-color-error': 'var(--mantine-color-error-10)', // used as error color
|
||||
'--mantine-color-placeholder': 'var(--mantine-color-secondary-10)', // used as placeholder color
|
||||
'--mantine-color-anchor': 'var(--mantine-color-secondary-10)', // used as anchor color
|
||||
|
||||
'--mantine-color-default': 'var(--mantine-color-secondary-0)', // used as default surface color
|
||||
'--mantine-color-default-hover': 'var(--mantine-color-secondary-1)', // used as default hover color
|
||||
'--mantine-color-default-color': 'var(--mantine-color-secondary-9)', // used as default text color
|
||||
'--mantine-color-default-border': 'var(--mantine-color-secondary-2)', // used as default border color
|
||||
'--mantine-color-dimmed': 'var(--mantine-color-secondary-10)', // used as dimmed text color
|
||||
|
||||
'--mantine-color-secondary-filled': 'var(--mantine-color-white)', // used as secondary surface color
|
||||
'--mantine-color-secondary-filled-hover': 'var(--mantine-color-secondary-1)', // used as secondary hover color
|
||||
|
||||
'--mantine-color-secondary-light': 'var(--mantine-color-secondary-1)', // used as primary light color
|
||||
'--mantine-color-secondary-light-hover': alpha('var(--mantine-color-secondary-light)', 0.8), // used as primary light hover color
|
||||
|
||||
'--mantine-color-secondary-text': 'var(--mantine-primary-color-contrast)', // can be used as secondary text color
|
||||
'--mantine-color-secondary-light-color': 'var(--mantine-color-secondary-8)', // used as primary light variant's text color
|
||||
|
||||
'--mantine-color-secondary-outline': 'var(--mantine-color-secondary-2)',
|
||||
'--mantine-color-secondary-outline-hover': 'var(--mantine-color-secondary-1)',
|
||||
|
||||
// all filled colors
|
||||
'--mantine-color-zinc-filled': 'var(--mantine-color-zinc-8)',
|
||||
'--mantine-color-zinc-filled-hover': alpha('var(--mantine-color-zinc-8)', 0.9),
|
||||
'--mantine-color-slate-filled': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-slate-filled-hover': alpha('var(--mantine-color-slate-8)', 0.9),
|
||||
'--mantine-color-gray-filled': 'var(--mantine-color-gray-8)',
|
||||
'--mantine-color-gray-filled-hover': alpha('var(--mantine-color-gray-8)', 0.9),
|
||||
'--mantine-color-neutral-filled': 'var(--mantine-color-neutral-8)',
|
||||
'--mantine-color-neutral-filled-hover': alpha('var(--mantine-color-neutral-8)', 0.9),
|
||||
'--mantine-color-stone-filled': 'var(--mantine-color-stone-8)',
|
||||
'--mantine-color-stone-filled-hover': alpha('var(--mantine-color-stone-8)', 0.9),
|
||||
'--mantine-color-red-filled': 'var(--mantine-color-red-5)',
|
||||
'--mantine-color-red-filled-hover': alpha('var(--mantine-color-red-5)', 0.9),
|
||||
'--mantine-color-rose-filled': 'var(--mantine-color-rose-5)',
|
||||
'--mantine-color-rose-filled-hover': alpha('var(--mantine-color-rose-5)', 0.9),
|
||||
'--mantine-color-orange-filled': 'var(--mantine-color-orange-5)',
|
||||
'--mantine-color-orange-filled-hover': alpha('var(--mantine-color-orange-5)', 0.9),
|
||||
'--mantine-color-amber-filled': 'var(--mantine-color-amber-5)',
|
||||
'--mantine-color-amber-filled-hover': alpha('var(--mantine-color-amber-5)', 0.9),
|
||||
'--mantine-color-yellow-filled': 'var(--mantine-color-yellow-4)',
|
||||
'--mantine-color-yellow-filled-hover': alpha('var(--mantine-color-yellow-4)', 0.9),
|
||||
'--mantine-color-lime-filled': 'var(--mantine-color-lime-5)',
|
||||
'--mantine-color-lime-filled-hover': alpha('var(--mantine-color-lime-5)', 0.9),
|
||||
'--mantine-color-green-filled': 'var(--mantine-color-green-6)',
|
||||
'--mantine-color-green-filled-hover': alpha('var(--mantine-color-green-6)', 0.9),
|
||||
'--mantine-color-emerald-filled': 'var(--mantine-color-emerald-5)',
|
||||
'--mantine-color-emerald-filled-hover': alpha('var(--mantine-color-emerald-5)', 0.9),
|
||||
'--mantine-color-teal-filled': 'var(--mantine-color-teal-5)',
|
||||
'--mantine-color-teal-filled-hover': alpha('var(--mantine-color-teal-5)', 0.9),
|
||||
'--mantine-color-cyan-filled': 'var(--mantine-color-cyan-5)',
|
||||
'--mantine-color-cyan-filled-hover': alpha('var(--mantine-color-cyan-5)', 0.9),
|
||||
'--mantine-color-sky-filled': 'var(--mantine-color-sky-5)',
|
||||
'--mantine-color-sky-filled-hover': alpha('var(--mantine-color-sky-5)', 0.9),
|
||||
'--mantine-color-blue-filled': 'var(--mantine-color-blue-6)',
|
||||
'--mantine-color-blue-filled-hover': alpha('var(--mantine-color-blue-6)', 0.9),
|
||||
'--mantine-color-indigo-filled': 'var(--mantine-color-indigo-5)',
|
||||
'--mantine-color-indigo-filled-hover': alpha('var(--mantine-color-indigo-5)', 0.9),
|
||||
'--mantine-color-violet-filled': 'var(--mantine-color-violet-5)',
|
||||
'--mantine-color-violet-filled-hover': alpha('var(--mantine-color-violet-5)', 0.9),
|
||||
'--mantine-color-purple-filled': 'var(--mantine-color-purple-5)',
|
||||
'--mantine-color-purple-filled-hover': alpha('var(--mantine-color-purple-5)', 0.9),
|
||||
'--mantine-color-fuchsia-filled': 'var(--mantine-color-fuchsia-5)',
|
||||
'--mantine-color-fuchsia-filled-hover': alpha('var(--mantine-color-fuchsia-5)', 0.9),
|
||||
'--mantine-color-pink-filled': 'var(--mantine-color-pink-5)',
|
||||
'--mantine-color-pink-filled-hover': alpha('var(--mantine-color-pink-5)', 0.9),
|
||||
|
||||
// all light colors
|
||||
'--mantine-color-zinc-light': alpha('var(--mantine-color-zinc-4)', 0.1),
|
||||
'--mantine-color-zinc-light-hover': alpha('var(--mantine-color-zinc-light)', 0.8),
|
||||
'--mantine-color-zinc-light-color': 'var(--mantine-color-zinc-6)',
|
||||
'--mantine-color-slate-light': alpha('var(--mantine-color-slate-4)', 0.1),
|
||||
'--mantine-color-slate-light-hover': alpha('var(--mantine-color-slate-light)', 0.8),
|
||||
'--mantine-color-slate-light-color': 'var(--mantine-color-slate-6)',
|
||||
'--mantine-color-gray-light': alpha('var(--mantine-color-gray-4)', 0.1),
|
||||
'--mantine-color-gray-light-hover': alpha('var(--mantine-color-gray-light)', 0.8),
|
||||
'--mantine-color-gray-light-color': 'var(--mantine-color-gray-6)',
|
||||
'--mantine-color-neutral-light': alpha('var(--mantine-color-neutral-4)', 0.1),
|
||||
'--mantine-color-neutral-light-hover': alpha('var(--mantine-color-neutral-light)', 0.8),
|
||||
'--mantine-color-neutral-light-color': 'var(--mantine-color-neutral-6)',
|
||||
'--mantine-color-stone-light': alpha('var(--mantine-color-stone-4)', 0.1),
|
||||
'--mantine-color-stone-light-hover': alpha('var(--mantine-color-stone-light)', 0.8),
|
||||
'--mantine-color-stone-light-color': 'var(--mantine-color-stone-6)',
|
||||
'--mantine-color-red-light': alpha('var(--mantine-color-red-4)', 0.1),
|
||||
'--mantine-color-red-light-hover': alpha('var(--mantine-color-red-light)', 0.8),
|
||||
'--mantine-color-red-light-color': 'var(--mantine-color-red-6)',
|
||||
'--mantine-color-rose-light': alpha('var(--mantine-color-rose-4)', 0.1),
|
||||
'--mantine-color-rose-light-hover': alpha('var(--mantine-color-rose-light)', 0.8),
|
||||
'--mantine-color-rose-light-color': 'var(--mantine-color-rose-6)',
|
||||
'--mantine-color-orange-light': alpha('var(--mantine-color-orange-4)', 0.1),
|
||||
'--mantine-color-orange-light-hover': alpha('var(--mantine-color-orange-light)', 0.8),
|
||||
'--mantine-color-orange-light-color': 'var(--mantine-color-orange-6)',
|
||||
'--mantine-color-amber-light': alpha('var(--mantine-color-amber-4)', 0.1),
|
||||
'--mantine-color-amber-light-hover': alpha('var(--mantine-color-amber-light)', 0.8),
|
||||
'--mantine-color-amber-light-color': 'var(--mantine-color-amber-6)',
|
||||
'--mantine-color-yellow-light': alpha('var(--mantine-color-yellow-4)', 0.1),
|
||||
'--mantine-color-yellow-light-hover': alpha('var(--mantine-color-yellow-light)', 0.8),
|
||||
'--mantine-color-yellow-light-color': 'var(--mantine-color-yellow-6)',
|
||||
'--mantine-color-lime-light': alpha('var(--mantine-color-lime-4)', 0.1),
|
||||
'--mantine-color-lime-light-hover': alpha('var(--mantine-color-lime-light)', 0.8),
|
||||
'--mantine-color-lime-light-color': 'var(--mantine-color-lime-6)',
|
||||
'--mantine-color-green-light': alpha('var(--mantine-color-green-4)', 0.1),
|
||||
'--mantine-color-green-light-hover': alpha('var(--mantine-color-green-light)', 0.8),
|
||||
'--mantine-color-green-light-color': 'var(--mantine-color-green-6)',
|
||||
'--mantine-color-emerald-light': alpha('var(--mantine-color-emerald-4)', 0.1),
|
||||
'--mantine-color-emerald-light-hover': alpha('var(--mantine-color-emerald-light)', 0.8),
|
||||
'--mantine-color-emerald-light-color': 'var(--mantine-color-emerald-6)',
|
||||
'--mantine-color-teal-light': alpha('var(--mantine-color-teal-4)', 0.1),
|
||||
'--mantine-color-teal-light-hover': alpha('var(--mantine-color-teal-light)', 0.8),
|
||||
'--mantine-color-teal-light-color': 'var(--mantine-color-teal-6)',
|
||||
'--mantine-color-cyan-light': alpha('var(--mantine-color-cyan-4)', 0.1),
|
||||
'--mantine-color-cyan-light-hover': alpha('var(--mantine-color-cyan-light)', 0.8),
|
||||
'--mantine-color-cyan-light-color': 'var(--mantine-color-cyan-6)',
|
||||
'--mantine-color-sky-light': alpha('var(--mantine-color-sky-4)', 0.1),
|
||||
'--mantine-color-sky-light-hover': alpha('var(--mantine-color-sky-light)', 0.8),
|
||||
'--mantine-color-sky-light-color': 'var(--mantine-color-sky-6)',
|
||||
'--mantine-color-blue-light': alpha('var(--mantine-color-blue-4)', 0.1),
|
||||
'--mantine-color-blue-light-hover': alpha('var(--mantine-color-blue-light)', 0.8),
|
||||
'--mantine-color-blue-light-color': 'var(--mantine-color-blue-6)',
|
||||
'--mantine-color-indigo-light': alpha('var(--mantine-color-indigo-4)', 0.1),
|
||||
'--mantine-color-indigo-light-hover': alpha('var(--mantine-color-indigo-light)', 0.8),
|
||||
'--mantine-color-indigo-light-color': 'var(--mantine-color-indigo-6)',
|
||||
'--mantine-color-violet-light': alpha('var(--mantine-color-violet-4)', 0.1),
|
||||
'--mantine-color-violet-light-hover': alpha('var(--mantine-color-violet-light)', 0.8),
|
||||
'--mantine-color-violet-light-color': 'var(--mantine-color-violet-6)',
|
||||
'--mantine-color-purple-light': alpha('var(--mantine-color-purple-4)', 0.1),
|
||||
'--mantine-color-purple-light-hover': alpha('var(--mantine-color-purple-light)', 0.8),
|
||||
'--mantine-color-purple-light-color': 'var(--mantine-color-purple-6)',
|
||||
'--mantine-color-fuchsia-light': alpha('var(--mantine-color-fuchsia-4)', 0.1),
|
||||
'--mantine-color-fuchsia-light-hover': alpha('var(--mantine-color-fuchsia-light)', 0.8),
|
||||
'--mantine-color-fuchsia-light-color': 'var(--mantine-color-fuchsia-6)',
|
||||
'--mantine-color-pink-light': alpha('var(--mantine-color-pink-4)', 0.1),
|
||||
'--mantine-color-pink-light-hover': alpha('var(--mantine-color-pink-light)', 0.8),
|
||||
'--mantine-color-pink-light-color': 'var(--mantine-color-pink-6)',
|
||||
|
||||
// all outline colors
|
||||
'--mantine-color-zinc-outline': 'var(--mantine-color-zinc-8)',
|
||||
'--mantine-color-zinc-outline-hover': alpha('var(--mantine-color-zinc-4)', 0.1),
|
||||
'--mantine-color-slate-outline': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-slate-outline-hover': alpha('var(--mantine-color-slate-4)', 0.1),
|
||||
'--mantine-color-gray-outline': 'var(--mantine-color-gray-8)',
|
||||
'--mantine-color-gray-outline-hover': alpha('var(--mantine-color-gray-4)', 0.1),
|
||||
'--mantine-color-neutral-outline': 'var(--mantine-color-neutral-8)',
|
||||
'--mantine-color-neutral-outline-hover': alpha('var(--mantine-color-neutral-4)', 0.1),
|
||||
'--mantine-color-stone-outline': 'var(--mantine-color-stone-8)',
|
||||
'--mantine-color-stone-outline-hover': alpha('var(--mantine-color-stone-4)', 0.1),
|
||||
'--mantine-color-red-outline': 'var(--mantine-color-red-5)',
|
||||
'--mantine-color-red-outline-hover': alpha('var(--mantine-color-red-4)', 0.1),
|
||||
'--mantine-color-rose-outline': 'var(--mantine-color-rose-5)',
|
||||
'--mantine-color-rose-outline-hover': alpha('var(--mantine-color-rose-4)', 0.1),
|
||||
'--mantine-color-orange-outline': 'var(--mantine-color-orange-5)',
|
||||
'--mantine-color-orange-outline-hover': alpha('var(--mantine-color-orange-4)', 0.1),
|
||||
'--mantine-color-amber-outline': 'var(--mantine-color-amber-5)',
|
||||
'--mantine-color-amber-outline-hover': alpha('var(--mantine-color-amber-4)', 0.1),
|
||||
'--mantine-color-yellow-outline': 'var(--mantine-color-yellow-4)',
|
||||
'--mantine-color-yellow-outline-hover': alpha('var(--mantine-color-yellow-4)', 0.1),
|
||||
'--mantine-color-lime-outline': 'var(--mantine-color-lime-5)',
|
||||
'--mantine-color-lime-outline-hover': alpha('var(--mantine-color-lime-4)', 0.1),
|
||||
'--mantine-color-green-outline': 'var(--mantine-color-green-6)',
|
||||
'--mantine-color-green-outline-hover': alpha('var(--mantine-color-green-4)', 0.1),
|
||||
'--mantine-color-emerald-outline': 'var(--mantine-color-emerald-5)',
|
||||
'--mantine-color-emerald-outline-hover': alpha('var(--mantine-color-emerald-4)', 0.1),
|
||||
'--mantine-color-teal-outline': 'var(--mantine-color-teal-5)',
|
||||
'--mantine-color-teal-outline-hover': alpha('var(--mantine-color-teal-4)', 0.1),
|
||||
'--mantine-color-cyan-outline': 'var(--mantine-color-cyan-5)',
|
||||
'--mantine-color-cyan-outline-hover': alpha('var(--mantine-color-cyan-4)', 0.1),
|
||||
'--mantine-color-sky-outline': 'var(--mantine-color-sky-5)',
|
||||
'--mantine-color-sky-outline-hover': alpha('var(--mantine-color-sky-4)', 0.1),
|
||||
'--mantine-color-blue-outline': 'var(--mantine-color-blue-6)',
|
||||
'--mantine-color-blue-outline-hover': alpha('var(--mantine-color-blue-4)', 0.1),
|
||||
'--mantine-color-indigo-outline': 'var(--mantine-color-indigo-5)',
|
||||
'--mantine-color-indigo-outline-hover': alpha('var(--mantine-color-indigo-4)', 0.1),
|
||||
'--mantine-color-violet-outline': 'var(--mantine-color-violet-5)',
|
||||
'--mantine-color-violet-outline-hover': alpha('var(--mantine-color-violet-4)', 0.1),
|
||||
'--mantine-color-purple-outline': 'var(--mantine-color-purple-5)',
|
||||
'--mantine-color-purple-outline-hover': alpha('var(--mantine-color-purple-4)', 0.1),
|
||||
'--mantine-color-fuchsia-outline': 'var(--mantine-color-fuchsia-5)',
|
||||
'--mantine-color-fuchsia-outline-hover': alpha('var(--mantine-color-fuchsia-4)', 0.1),
|
||||
'--mantine-color-pink-outline': 'var(--mantine-color-pink-5)',
|
||||
'--mantine-color-pink-outline-hover': alpha('var(--mantine-color-pink-4)', 0.1),
|
||||
|
||||
// all contrast colors
|
||||
'--mantine-color-zinc-contrast': 'var(--mantine-color-zinc-0)',
|
||||
'--mantine-color-slate-contrast': 'var(--mantine-color-slate-0)',
|
||||
'--mantine-color-gray-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-neutral-contrast': 'var(--mantine-color-neutral-0)',
|
||||
'--mantine-color-stone-contrast': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-red-contrast': 'var(--mantine-color-red-0)',
|
||||
'--mantine-color-rose-contrast': 'var(--mantine-color-rose-0)',
|
||||
'--mantine-color-orange-contrast': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-amber-contrast': 'var(--mantine-color-amber-0)',
|
||||
'--mantine-color-yellow-contrast': '#422006',
|
||||
'--mantine-color-lime-contrast': 'var(--mantine-color-lime-0)',
|
||||
'--mantine-color-green-contrast': 'var(--mantine-color-rose-0)',
|
||||
'--mantine-color-emerald-contrast': 'var(--mantine-color-emerald-0)',
|
||||
'--mantine-color-teal-contrast': 'var(--mantine-color-teal-0)',
|
||||
'--mantine-color-cyan-contrast': 'var(--mantine-color-cyan-0)',
|
||||
'--mantine-color-sky-contrast': 'var(--mantine-color-sky-0)',
|
||||
'--mantine-color-blue-contrast': 'var(--mantine-color-slate-0)',
|
||||
'--mantine-color-indigo-contrast': 'var(--mantine-color-indigo-0)',
|
||||
'--mantine-color-violet-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-purple-contrast': 'var(--mantine-color-purple-0)',
|
||||
'--mantine-color-fuchsia-contrast': 'var(--mantine-color-fuchsia-0)',
|
||||
'--mantine-color-pink-contrast': 'var(--mantine-color-pink-0)',
|
||||
},
|
||||
dark: {
|
||||
// all variables that depend on dark color scheme
|
||||
'--mantine-primary-color-contrast': 'var(--mantine-color-zinc-8)', // used as primary color contrast
|
||||
'--mantine-color-text': 'var(--mantine-color-secondary-0)', // used as text color
|
||||
'--mantine-color-body': 'var(--mantine-color-secondary-9)', // used as body color
|
||||
'--mantine-color-error': 'var(--mantine-color-error-10)', // used as error color
|
||||
'--mantine-color-placeholder': 'var(--mantine-color-secondary-4)', // used as placeholder color
|
||||
'--mantine-color-anchor': 'var(--mantine-color-secondary-4)', // used as anchor color
|
||||
|
||||
'--mantine-color-default': 'var(--mantine-color-secondary-9)', // used as default surface color
|
||||
'--mantine-color-default-hover': 'var(--mantine-color-secondary-7)', // used as default hover color
|
||||
'--mantine-color-default-color': 'var(--mantine-color-secondary-1)', // used as default text color
|
||||
'--mantine-color-default-border': 'var(--mantine-color-secondary-7)', // used as default border color
|
||||
'--mantine-color-dimmed': 'var(--mantine-color-secondary-4)', // used as dimmed text color
|
||||
|
||||
'--mantine-color-secondary-filled': 'var(--mantine-color-secondary-8)', // used as secondary surface color
|
||||
'--mantine-color-secondary-filled-hover': alpha('var(--mantine-color-secondary-filled)', 0.9), // used as secondary hover color
|
||||
|
||||
'--mantine-color-secondary-light': 'var(--mantine-color-secondary-7)', // used as primary light color
|
||||
'--mantine-color-secondary-light-hover': alpha('var(--mantine-color-secondary-light)', 0.8), // used as primary light hover color
|
||||
|
||||
'--mantine-color-secondary-text': 'var(--mantine-primary-color-contrast)', // can be used as secondary text color
|
||||
'--mantine-color-secondary-light-color': 'var(--mantine-color-secondary-0)', // used as primary light text color
|
||||
|
||||
'--mantine-color-secondary-outline': 'var(--mantine-color-secondary-7)',
|
||||
'--mantine-color-secondary-outline-hover': 'var(--mantine-color-secondary-7)',
|
||||
|
||||
// all filled colors
|
||||
'--mantine-color-zinc-filled': 'var(--mantine-color-zinc-0)',
|
||||
'--mantine-color-zinc-filled-hover': alpha('var(--mantine-color-zinc-0)', 0.9),
|
||||
'--mantine-color-slate-filled': 'var(--mantine-color-slate-0)',
|
||||
'--mantine-color-slate-filled-hover': alpha('var(--mantine-color-slate-0)', 0.9),
|
||||
'--mantine-color-gray-filled': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-gray-filled-hover': alpha('var(--mantine-color-gray-0)', 0.9),
|
||||
'--mantine-color-neutral-filled': 'var(--mantine-color-neutral-0)',
|
||||
'--mantine-color-neutral-filled-hover': alpha('var(--mantine-color-neutral-0)', 0.9),
|
||||
'--mantine-color-stone-filled': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-stone-filled-hover': alpha('var(--mantine-color-stone-0)', 0.9),
|
||||
'--mantine-color-red-filled': 'var(--mantine-color-red-5)',
|
||||
'--mantine-color-red-filled-hover': alpha('var(--mantine-color-red-5)', 0.9),
|
||||
'--mantine-color-rose-filled': 'var(--mantine-color-rose-5)',
|
||||
'--mantine-color-rose-filled-hover': alpha('var(--mantine-color-rose-5)', 0.9),
|
||||
'--mantine-color-orange-filled': 'var(--mantine-color-orange-6)',
|
||||
'--mantine-color-orange-filled-hover': alpha('var(--mantine-color-orange-6)', 0.9),
|
||||
'--mantine-color-amber-filled': 'var(--mantine-color-amber-5)',
|
||||
'--mantine-color-amber-filled-hover': alpha('var(--mantine-color-amber-5)', 0.9),
|
||||
'--mantine-color-yellow-filled': 'var(--mantine-color-yellow-4)',
|
||||
'--mantine-color-yellow-filled-hover': alpha('var(--mantine-color-yellow-4)', 0.9),
|
||||
'--mantine-color-lime-filled': 'var(--mantine-color-lime-4)',
|
||||
'--mantine-color-lime-filled-hover': alpha('var(--mantine-color-lime-4)', 0.9),
|
||||
'--mantine-color-green-filled': 'var(--mantine-color-green-5)',
|
||||
'--mantine-color-green-filled-hover': alpha('var(--mantine-color-green-5)', 0.9),
|
||||
'--mantine-color-emerald-filled': 'var(--mantine-color-emerald-5)',
|
||||
'--mantine-color-emerald-filled-hover': alpha('var(--mantine-color-emerald-5)', 0.9),
|
||||
'--mantine-color-teal-filled': 'var(--mantine-color-teal-4)',
|
||||
'--mantine-color-teal-filled-hover': alpha('var(--mantine-color-teal-4)', 0.9),
|
||||
'--mantine-color-cyan-filled': 'var(--mantine-color-cyan-4)',
|
||||
'--mantine-color-cyan-filled-hover': alpha('var(--mantine-color-cyan-4)', 0.9),
|
||||
'--mantine-color-sky-filled': 'var(--mantine-color-sky-4)',
|
||||
'--mantine-color-sky-filled-hover': alpha('var(--mantine-color-sky-4)', 0.9),
|
||||
'--mantine-color-blue-filled': 'var(--mantine-color-blue-5)',
|
||||
'--mantine-color-blue-filled-hover': alpha('var(--mantine-color-blue-5)', 0.9),
|
||||
'--mantine-color-indigo-filled': 'var(--mantine-color-indigo-6)',
|
||||
'--mantine-color-indigo-filled-hover': alpha('var(--mantine-color-indigo-6)', 0.9),
|
||||
'--mantine-color-violet-filled': 'var(--mantine-color-violet-6)',
|
||||
'--mantine-color-violet-filled-hover': alpha('var(--mantine-color-violet-6)', 0.9),
|
||||
'--mantine-color-purple-filled': 'var(--mantine-color-purple-6)',
|
||||
'--mantine-color-purple-filled-hover': alpha('var(--mantine-color-purple-6)', 0.9),
|
||||
'--mantine-color-fuchsia-filled': 'var(--mantine-color-fuchsia-7)',
|
||||
'--mantine-color-fuchsia-filled-hover': alpha('var(--mantine-color-fuchsia-7)', 0.9),
|
||||
'--mantine-color-pink-filled': 'var(--mantine-color-pink-6)',
|
||||
'--mantine-color-pink-filled-hover': alpha('var(--mantine-color-pink-6)', 0.9),
|
||||
|
||||
// all light colors
|
||||
'--mantine-color-zinc-light': alpha('var(--mantine-color-zinc-4)', 0.15),
|
||||
'--mantine-color-zinc-light-hover': alpha('var(--mantine-color-zinc-light)', 0.8),
|
||||
'--mantine-color-zinc-light-color': 'var(--mantine-color-zinc-3)',
|
||||
'--mantine-color-slate-light': alpha('var(--mantine-color-slate-4)', 0.15),
|
||||
'--mantine-color-slate-light-hover': alpha('var(--mantine-color-slate-light)', 0.8),
|
||||
'--mantine-color-slate-light-color': 'var(--mantine-color-slate-3)',
|
||||
'--mantine-color-gray-light': alpha('var(--mantine-color-gray-4)', 0.15),
|
||||
'--mantine-color-gray-light-hover': alpha('var(--mantine-color-gray-light)', 0.8),
|
||||
'--mantine-color-gray-light-color': 'var(--mantine-color-gray-3)',
|
||||
'--mantine-color-neutral-light': alpha('var(--mantine-color-neutral-4)', 0.15),
|
||||
'--mantine-color-neutral-light-hover': alpha('var(--mantine-color-neutral-light)', 0.8),
|
||||
'--mantine-color-neutral-light-color': 'var(--mantine-color-neutral-3)',
|
||||
'--mantine-color-stone-light': alpha('var(--mantine-color-stone-4)', 0.15),
|
||||
'--mantine-color-stone-light-hover': alpha('var(--mantine-color-stone-light)', 0.8),
|
||||
'--mantine-color-stone-light-color': 'var(--mantine-color-stone-3)',
|
||||
'--mantine-color-red-light': alpha('var(--mantine-color-red-4)', 0.15),
|
||||
'--mantine-color-red-light-hover': alpha('var(--mantine-color-red-light)', 0.8),
|
||||
'--mantine-color-red-light-color': 'var(--mantine-color-red-3)',
|
||||
'--mantine-color-rose-light': alpha('var(--mantine-color-rose-4)', 0.15),
|
||||
'--mantine-color-rose-light-hover': alpha('var(--mantine-color-rose-light)', 0.8),
|
||||
'--mantine-color-rose-light-color': 'var(--mantine-color-rose-3)',
|
||||
'--mantine-color-orange-light': alpha('var(--mantine-color-orange-4)', 0.15),
|
||||
'--mantine-color-orange-light-hover': alpha('var(--mantine-color-orange-light)', 0.8),
|
||||
'--mantine-color-orange-light-color': 'var(--mantine-color-orange-3)',
|
||||
'--mantine-color-amber-light': alpha('var(--mantine-color-amber-4)', 0.15),
|
||||
'--mantine-color-amber-light-hover': alpha('var(--mantine-color-amber-light)', 0.8),
|
||||
'--mantine-color-amber-light-color': 'var(--mantine-color-amber-3)',
|
||||
'--mantine-color-yellow-light': alpha('var(--mantine-color-yellow-4)', 0.15),
|
||||
'--mantine-color-yellow-light-hover': alpha('var(--mantine-color-yellow-light)', 0.8),
|
||||
'--mantine-color-yellow-light-color': 'var(--mantine-color-yellow-3)',
|
||||
'--mantine-color-lime-light': alpha('var(--mantine-color-lime-4)', 0.15),
|
||||
'--mantine-color-lime-light-hover': alpha('var(--mantine-color-lime-light)', 0.8),
|
||||
'--mantine-color-lime-light-color': 'var(--mantine-color-lime-3)',
|
||||
'--mantine-color-green-light': alpha('var(--mantine-color-green-4)', 0.15),
|
||||
'--mantine-color-green-light-hover': alpha('var(--mantine-color-green-light)', 0.8),
|
||||
'--mantine-color-green-light-color': 'var(--mantine-color-green-3)',
|
||||
'--mantine-color-emerald-light': alpha('var(--mantine-color-emerald-4)', 0.15),
|
||||
'--mantine-color-emerald-light-hover': alpha('var(--mantine-color-emerald-light)', 0.8),
|
||||
'--mantine-color-emerald-light-color': 'var(--mantine-color-emerald-3)',
|
||||
'--mantine-color-teal-light': alpha('var(--mantine-color-teal-4)', 0.15),
|
||||
'--mantine-color-teal-light-hover': alpha('var(--mantine-color-teal-light)', 0.8),
|
||||
'--mantine-color-teal-light-color': 'var(--mantine-color-teal-3)',
|
||||
'--mantine-color-cyan-light': alpha('var(--mantine-color-cyan-4)', 0.15),
|
||||
'--mantine-color-cyan-light-hover': alpha('var(--mantine-color-cyan-light)', 0.8),
|
||||
'--mantine-color-cyan-light-color': 'var(--mantine-color-cyan-3)',
|
||||
'--mantine-color-sky-light': alpha('var(--mantine-color-sky-4)', 0.15),
|
||||
'--mantine-color-sky-light-hover': alpha('var(--mantine-color-sky-light)', 0.8),
|
||||
'--mantine-color-sky-light-color': 'var(--mantine-color-sky-3)',
|
||||
'--mantine-color-blue-light': alpha('var(--mantine-color-blue-4)', 0.15),
|
||||
'--mantine-color-blue-light-hover': alpha('var(--mantine-color-blue-light)', 0.8),
|
||||
'--mantine-color-blue-light-color': 'var(--mantine-color-blue-3)',
|
||||
'--mantine-color-indigo-light': alpha('var(--mantine-color-indigo-4)', 0.15),
|
||||
'--mantine-color-indigo-light-hover': alpha('var(--mantine-color-indigo-light)', 0.8),
|
||||
'--mantine-color-indigo-light-color': 'var(--mantine-color-indigo-3)',
|
||||
'--mantine-color-violet-light': alpha('var(--mantine-color-violet-4)', 0.15),
|
||||
'--mantine-color-violet-light-hover': alpha('var(--mantine-color-violet-light)', 0.8),
|
||||
'--mantine-color-violet-light-color': 'var(--mantine-color-violet-3)',
|
||||
'--mantine-color-purple-light': alpha('var(--mantine-color-purple-4)', 0.15),
|
||||
'--mantine-color-purple-light-hover': alpha('var(--mantine-color-purple-light)', 0.8),
|
||||
'--mantine-color-purple-light-color': 'var(--mantine-color-purple-3)',
|
||||
'--mantine-color-fuchsia-light': alpha('var(--mantine-color-fuchsia-4)', 0.15),
|
||||
'--mantine-color-fuchsia-light-hover': alpha('var(--mantine-color-fuchsia-light)', 0.8),
|
||||
'--mantine-color-fuchsia-light-color': 'var(--mantine-color-fuchsia-3)',
|
||||
'--mantine-color-pink-light': alpha('var(--mantine-color-pink-4)', 0.15),
|
||||
'--mantine-color-pink-light-hover': alpha('var(--mantine-color-pink-light)', 0.8),
|
||||
'--mantine-color-pink-light-color': 'var(--mantine-color-pink-3)',
|
||||
|
||||
// all outline colors
|
||||
'--mantine-color-zinc-outline': 'var(--mantine-color-zinc-0)',
|
||||
'--mantine-color-zinc-outline-hover': alpha('var(--mantine-color-zinc-4)', 0.15),
|
||||
'--mantine-color-slate-outline': 'var(--mantine-color-slate-0)',
|
||||
'--mantine-color-slate-outline-hover': alpha('var(--mantine-color-slate-4)', 0.15),
|
||||
'--mantine-color-gray-outline': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-gray-outline-hover': alpha('var(--mantine-color-gray-4)', 0.15),
|
||||
'--mantine-color-neutral-outline': 'var(--mantine-color-neutral-0)',
|
||||
'--mantine-color-neutral-outline-hover': alpha('var(--mantine-color-neutral-4)', 0.15),
|
||||
'--mantine-color-stone-outline': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-stone-outline-hover': alpha('var(--mantine-color-stone-4)', 0.15),
|
||||
'--mantine-color-red-outline': 'var(--mantine-color-red-5)',
|
||||
'--mantine-color-red-outline-hover': alpha('var(--mantine-color-red-4)', 0.15),
|
||||
'--mantine-color-rose-outline': 'var(--mantine-color-rose-5)',
|
||||
'--mantine-color-rose-outline-hover': alpha('var(--mantine-color-rose-4)', 0.15),
|
||||
'--mantine-color-orange-outline': 'var(--mantine-color-orange-6)',
|
||||
'--mantine-color-orange-outline-hover': alpha('var(--mantine-color-orange-4)', 0.15),
|
||||
'--mantine-color-amber-outline': 'var(--mantine-color-amber-5)',
|
||||
'--mantine-color-amber-outline-hover': alpha('var(--mantine-color-amber-4)', 0.15),
|
||||
'--mantine-color-yellow-outline': 'var(--mantine-color-yellow-4)',
|
||||
'--mantine-color-yellow-outline-hover': alpha('var(--mantine-color-yellow-4)', 0.15),
|
||||
'--mantine-color-lime-outline': 'var(--mantine-color-lime-4)',
|
||||
'--mantine-color-lime-outline-hover': alpha('var(--mantine-color-lime-4)', 0.15),
|
||||
'--mantine-color-green-outline': 'var(--mantine-color-green-5)',
|
||||
'--mantine-color-green-outline-hover': alpha('var(--mantine-color-green-4)', 0.15),
|
||||
'--mantine-color-emerald-outline': 'var(--mantine-color-emerald-5)',
|
||||
'--mantine-color-emerald-outline-hover': alpha('var(--mantine-color-emerald-4)', 0.15),
|
||||
'--mantine-color-teal-outline': 'var(--mantine-color-teal-4)',
|
||||
'--mantine-color-teal-outline-hover': alpha('var(--mantine-color-teal-4)', 0.15),
|
||||
'--mantine-color-cyan-outline': 'var(--mantine-color-cyan-4)',
|
||||
'--mantine-color-cyan-outline-hover': alpha('var(--mantine-color-cyan-4)', 0.15),
|
||||
'--mantine-color-sky-outline': 'var(--mantine-color-sky-4)',
|
||||
'--mantine-color-sky-outline-hover': alpha('var(--mantine-color-sky-4)', 0.15),
|
||||
'--mantine-color-blue-outline': 'var(--mantine-color-blue-5)',
|
||||
'--mantine-color-blue-outline-hover': alpha('var(--mantine-color-blue-4)', 0.15),
|
||||
'--mantine-color-indigo-outline': 'var(--mantine-color-indigo-6)',
|
||||
'--mantine-color-indigo-outline-hover': alpha('var(--mantine-color-indigo-4)', 0.15),
|
||||
'--mantine-color-violet-outline': 'var(--mantine-color-violet-6)',
|
||||
'--mantine-color-violet-outline-hover': alpha('var(--mantine-color-violet-4)', 0.15),
|
||||
'--mantine-color-purple-outline': 'var(--mantine-color-purple-6)',
|
||||
'--mantine-color-purple-outline-hover': alpha('var(--mantine-color-purple-4)', 0.15),
|
||||
'--mantine-color-fuchsia-outline': 'var(--mantine-color-fuchsia-7)',
|
||||
'--mantine-color-fuchsia-outline-hover': alpha('var(--mantine-color-fuchsia-4)', 0.15),
|
||||
'--mantine-color-pink-outline': 'var(--mantine-color-pink-6)',
|
||||
'--mantine-color-pink-outline-hover': alpha('var(--mantine-color-pink-4)', 0.15),
|
||||
|
||||
// all contrast colors
|
||||
'--mantine-color-zinc-contrast': 'var(--mantine-color-zinc-8)',
|
||||
'--mantine-color-slate-contrast': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-gray-contrast': 'var(--mantine-color-gray-8)',
|
||||
'--mantine-color-neutral-contrast': 'var(--mantine-color-neutral-8)',
|
||||
'--mantine-color-stone-contrast': 'var(--mantine-color-stone-8)',
|
||||
'--mantine-color-red-contrast': 'var(--mantine-color-red-0)',
|
||||
'--mantine-color-rose-contrast': 'var(--mantine-color-rose-0)',
|
||||
'--mantine-color-orange-contrast': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-amber-contrast': 'var(--mantine-color-stone-8)',
|
||||
'--mantine-color-yellow-contrast': '#422006',
|
||||
'--mantine-color-lime-contrast': 'var(--mantine-color-stone-8)',
|
||||
'--mantine-color-green-contrast': 'var(--mantine-color-green-9)',
|
||||
'--mantine-color-emerald-contrast': 'var(--mantine-color-stone-0)',
|
||||
'--mantine-color-teal-contrast': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-cyan-contrast': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-sky-contrast': 'var(--mantine-color-slate-8)',
|
||||
'--mantine-color-blue-contrast': 'var(--mantine-color-slate-0)',
|
||||
'--mantine-color-indigo-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-violet-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-purple-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-fuchsia-contrast': 'var(--mantine-color-gray-0)',
|
||||
'--mantine-color-pink-contrast': 'var(--mantine-color-gray-0)',
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1,4 @@
|
||||
<svg width="16" height="16" viewBox="0 0 16 16" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M8.06935 0.740967C8.46471 0.740967 8.78513 1.06143 8.78513 1.45675C8.78513 1.73028 8.6317 1.96783 8.40619 2.08833V2.6357C8.60265 2.66378 9.18471 2.67955 9.92197 3.01465C10.8483 3.4357 11.121 3.77803 11.4378 4.2357C11.6095 4.48376 11.7205 4.74914 11.7909 5.0357H11.9009C12.273 5.0357 12.5746 5.33732 12.5746 5.70938V6.38307C12.5746 6.75514 12.273 7.05675 11.9009 7.05675H11.8192C11.7591 7.40026 11.6762 7.69328 11.6062 7.85675C11.413 8.30755 11.129 8.48833 10.9746 8.53044C11.0869 8.57254 11.4395 8.67604 11.6483 8.82517C11.943 9.0357 12.2378 9.39174 12.2378 9.75149C12.2378 10.0462 12.1957 10.4673 11.9851 10.6778C11.7074 10.9556 11.2272 11.4778 10.9746 11.6883L6.34303 15.8568L7.35356 13.4989L10.3851 9.54096H8.02724L8.86934 6.59359L8.19658 7.34393L8.19567 7.35149L5.2483 10.762H7.6904L7.10093 12.7831L5.62724 11.8989L4.91146 11.4357C4.53251 11.1831 4.32198 11.0989 4.06935 10.6778C3.91617 10.4225 3.90093 10.0462 3.90093 9.75149C3.90093 9.39174 4.19567 9.0357 4.4904 8.82517C4.69915 8.67604 4.78514 8.61465 4.99567 8.53044C4.82724 8.44623 4.7257 8.30755 4.53251 7.85675C4.46245 7.69328 4.37956 7.40026 4.31951 7.05675H4.23777C3.8657 7.05675 3.56409 6.75514 3.56409 6.38307V5.70938C3.56409 5.33732 3.8657 5.0357 4.23777 5.0357H4.34788C4.41815 4.74914 4.5292 4.48376 4.70093 4.2357C5.01778 3.77803 5.2904 3.4357 6.21672 3.01465C6.95396 2.67955 7.53602 2.66378 7.73251 2.6357V2.08833C7.50704 1.96783 7.35356 1.73028 7.35356 1.45675C7.35356 1.06143 7.67403 0.740967 8.06935 0.740967ZM6.80619 5.0357C6.50389 5.0357 6.25882 5.28077 6.25882 5.58307C6.25882 5.88538 6.50389 6.13044 6.80619 6.13044C7.1085 6.13044 7.35356 5.88538 7.35356 5.58307C7.35356 5.28077 7.1085 5.0357 6.80619 5.0357ZM9.3325 5.0357C9.03018 5.0357 8.78513 5.28077 8.78513 5.58307C8.78513 5.88538 9.03018 6.13044 9.3325 6.13044C9.63481 6.13044 9.87987 5.88538 9.87987 5.58307C9.87987 5.28077 9.63481 5.0357 9.3325 5.0357Z" fill="#F69047"/>
|
||||
<path d="M12.2279 9.63738C12.2342 9.67527 12.2378 9.71342 12.2378 9.75165C12.2378 10.0464 12.1957 10.4674 11.9851 10.678C11.7074 10.9558 11.2272 11.478 10.9746 11.6885L6.34305 15.8569L7.35357 13.499L7.9831 12.677C8.20912 12.6273 8.41543 12.5774 8.57462 12.5306C9.29041 12.3201 10.1325 11.6885 10.7641 11.099C11.2418 10.6531 11.9076 9.97136 12.2279 9.63738ZM9.62725 3.77271C10.3248 3.77271 10.8904 4.33825 10.8904 5.03586V6.80428C10.8904 7.50191 10.3248 8.06744 9.62725 8.06744H8.4483L8.86935 6.59376L8.19659 7.34409L8.19568 7.35165L7.57479 8.06744H6.59568C5.89805 8.06744 5.33252 7.50191 5.33252 6.80428V5.03586C5.33252 4.33825 5.89805 3.77271 6.59568 3.77271H9.62725ZM6.8062 5.03586C6.5039 5.03586 6.25884 5.28093 6.25884 5.58323C6.25884 5.88554 6.5039 6.1306 6.8062 6.1306C7.10851 6.1306 7.35357 5.88554 7.35357 5.58323C7.35357 5.28093 7.10851 5.03586 6.8062 5.03586ZM9.33251 5.03586C9.0302 5.03586 8.78514 5.28093 8.78514 5.58323C8.78514 5.88554 9.0302 6.1306 9.33251 6.1306C9.63483 6.1306 9.87988 5.88554 9.87988 5.58323C9.87988 5.28093 9.63483 5.03586 9.33251 5.03586Z" fill="#C45259"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 3.0 KiB |
@@ -0,0 +1,4 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
export * from './slice';
|
||||
export * from './selectors';
|
||||
@@ -0,0 +1,8 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { RootState } from '../../store';
|
||||
|
||||
export const selectConfig = (state: RootState) => state.config;
|
||||
export const selectAutoRefreshMs = (state: RootState) => state.config.autoRefreshMs;
|
||||
export const selectBaseUrl = (state: RootState) => state.config.baseUrl;
|
||||
export const selectThemePreference = (state: RootState) => state.config.theme;
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSlice, PayloadAction } from '@reduxjs/toolkit';
|
||||
import type { ConfigState, ThemePreference } from '@/types';
|
||||
|
||||
export const initialConfigState: ConfigState = {
|
||||
baseUrl: typeof window !== 'undefined' ? window.location.origin : '',
|
||||
autoRefreshMs: 0,
|
||||
theme: 'system',
|
||||
};
|
||||
|
||||
const configSlice = createSlice({
|
||||
name: 'config',
|
||||
initialState: initialConfigState,
|
||||
reducers: {
|
||||
setBaseUrl(state, action: PayloadAction<string>) {
|
||||
state.baseUrl = action.payload;
|
||||
},
|
||||
setAutoRefreshMs(state, action: PayloadAction<number>) {
|
||||
state.autoRefreshMs = action.payload;
|
||||
},
|
||||
setTheme(state, action: PayloadAction<ThemePreference>) {
|
||||
state.theme = action.payload;
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
export const { setAutoRefreshMs, setBaseUrl, setTheme } = configSlice.actions;
|
||||
|
||||
export const configReducer = configSlice.reducer;
|
||||
@@ -0,0 +1,5 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
export * from './slice';
|
||||
export * from './selectors';
|
||||
export { useGetResourcesQuery } from '../rollouts';
|
||||
@@ -0,0 +1,94 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createServerBackedStore } from '@test-utils';
|
||||
import { describe, expect, it } from 'vitest';
|
||||
import { rolloutsApi } from '@/features/rollouts';
|
||||
import type { Resources } from '@/types';
|
||||
import { selectResourcesQueryArgs } from './selectors';
|
||||
import {
|
||||
resetResourcesFilters,
|
||||
setResourcesPage,
|
||||
setResourcesRecordsPerPage,
|
||||
setResourcesSearchTerm,
|
||||
setResourcesSort,
|
||||
} from './slice';
|
||||
|
||||
const extractResourceIds = (resources: Resources[]): string[] => resources.map((resource) => resource.resourcesId);
|
||||
|
||||
describe('resources feature integration', () => {
|
||||
it('builds default query arguments from the UI state', () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectResourcesQueryArgs(store.getState());
|
||||
|
||||
expect(queryArgs).toMatchObject({
|
||||
limit: 50,
|
||||
offset: 0,
|
||||
sortBy: 'update_time',
|
||||
sortOrder: 'desc',
|
||||
resourcesIdContains: undefined,
|
||||
});
|
||||
});
|
||||
|
||||
it('fetches resources from the Python LightningStore server', async () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectResourcesQueryArgs(store.getState());
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getResources.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.total).toBe(5);
|
||||
expect(data.items).toHaveLength(5);
|
||||
|
||||
const resourceIds = extractResourceIds(data.items);
|
||||
expect(resourceIds).toEqual(expect.arrayContaining(['rs-story-001', 'rs-story-005']));
|
||||
|
||||
const updateTimes = data.items.map((resource) => resource.updateTime);
|
||||
const sortedUpdateTimes = [...updateTimes].sort((a, b) => b - a);
|
||||
expect(updateTimes).toEqual(sortedUpdateTimes);
|
||||
|
||||
expect(data.items[0].resources).toBeDefined();
|
||||
expect(Object.keys(data.items[0].resources)).not.toHaveLength(0);
|
||||
});
|
||||
|
||||
it('paginates resource results based on UI state', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(setResourcesRecordsPerPage(2));
|
||||
store.dispatch(setResourcesPage(2));
|
||||
|
||||
const queryArgs = selectResourcesQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({ limit: 2, offset: 2 });
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getResources.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(2);
|
||||
expect(data.total).toBe(5);
|
||||
expect(data.items.map((resource) => resource.resourcesId)).toEqual(['rs-story-005', 'rs-story-002']);
|
||||
});
|
||||
|
||||
it('applies search and sorting preferences', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(resetResourcesFilters());
|
||||
store.dispatch(setResourcesSearchTerm('rs-story-003'));
|
||||
store.dispatch(setResourcesSort({ column: 'version', direction: 'asc' }));
|
||||
|
||||
const queryArgs = selectResourcesQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({
|
||||
limit: 50,
|
||||
offset: 0,
|
||||
sortBy: 'version',
|
||||
sortOrder: 'asc',
|
||||
resourcesIdContains: 'rs-story-003',
|
||||
});
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getResources.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(1);
|
||||
expect(data.items[0].resourcesId).toBe('rs-story-003');
|
||||
expect(data.items[0].version).toBeGreaterThan(0);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,41 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSelector } from '@reduxjs/toolkit';
|
||||
import type { GetResourcesQueryArgs } from '@/features/rollouts';
|
||||
import type { RootState } from '@/store';
|
||||
import type { ResourcesSortState } from './slice';
|
||||
|
||||
const RESOURCES_SORT_FIELD_MAP: Record<string, string> = {
|
||||
resourcesId: 'resources_id',
|
||||
version: 'version',
|
||||
createTime: 'create_time',
|
||||
updateTime: 'update_time',
|
||||
};
|
||||
|
||||
const resolveResourcesSortField = (sort: ResourcesSortState): string =>
|
||||
RESOURCES_SORT_FIELD_MAP[sort.column] ?? 'update_time';
|
||||
|
||||
export const selectResourcesUiState = (state: RootState) => state.resources;
|
||||
|
||||
export const selectResourcesSearchTerm = (state: RootState) => selectResourcesUiState(state).searchTerm;
|
||||
export const selectResourcesPage = (state: RootState) => selectResourcesUiState(state).page;
|
||||
export const selectResourcesRecordsPerPage = (state: RootState) => selectResourcesUiState(state).recordsPerPage;
|
||||
export const selectResourcesSort = (state: RootState) => selectResourcesUiState(state).sort;
|
||||
|
||||
export const selectResourcesQueryArgs = createSelector(
|
||||
[selectResourcesSearchTerm, selectResourcesPage, selectResourcesRecordsPerPage, selectResourcesSort],
|
||||
(searchTerm, page, recordsPerPage, sort): GetResourcesQueryArgs => {
|
||||
const normalizedSearch = searchTerm.trim();
|
||||
const limit = Math.max(1, recordsPerPage);
|
||||
const offset = Math.max(0, (page - 1) * limit);
|
||||
const sortBy = resolveResourcesSortField(sort);
|
||||
|
||||
return {
|
||||
limit,
|
||||
offset,
|
||||
sortBy,
|
||||
sortOrder: sort.direction,
|
||||
resourcesIdContains: normalizedSearch.length > 0 ? normalizedSearch : undefined,
|
||||
};
|
||||
},
|
||||
);
|
||||
@@ -0,0 +1,64 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createSlice, type PayloadAction } from '@reduxjs/toolkit';
|
||||
|
||||
export type SortDirection = 'asc' | 'desc';
|
||||
|
||||
export type ResourcesSortState = {
|
||||
column: string;
|
||||
direction: SortDirection;
|
||||
};
|
||||
|
||||
export type ResourcesUiState = {
|
||||
searchTerm: string;
|
||||
page: number;
|
||||
recordsPerPage: number;
|
||||
sort: ResourcesSortState;
|
||||
};
|
||||
|
||||
export const initialResourcesUiState: ResourcesUiState = {
|
||||
searchTerm: '',
|
||||
page: 1,
|
||||
recordsPerPage: 50,
|
||||
sort: {
|
||||
column: 'updateTime',
|
||||
direction: 'desc',
|
||||
},
|
||||
};
|
||||
|
||||
const resourcesSlice = createSlice({
|
||||
name: 'resources',
|
||||
initialState: initialResourcesUiState,
|
||||
reducers: {
|
||||
setResourcesSearchTerm(state, action: PayloadAction<string>) {
|
||||
state.searchTerm = action.payload;
|
||||
state.page = 1;
|
||||
},
|
||||
setResourcesPage(state, action: PayloadAction<number>) {
|
||||
state.page = action.payload;
|
||||
},
|
||||
setResourcesRecordsPerPage(state, action: PayloadAction<number>) {
|
||||
state.recordsPerPage = action.payload;
|
||||
state.page = 1;
|
||||
},
|
||||
setResourcesSort(state, action: PayloadAction<ResourcesSortState>) {
|
||||
state.sort = action.payload;
|
||||
},
|
||||
resetResourcesFilters(state) {
|
||||
state.searchTerm = initialResourcesUiState.searchTerm;
|
||||
state.page = initialResourcesUiState.page;
|
||||
state.recordsPerPage = initialResourcesUiState.recordsPerPage;
|
||||
state.sort = initialResourcesUiState.sort;
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
export const {
|
||||
setResourcesSearchTerm,
|
||||
setResourcesPage,
|
||||
setResourcesRecordsPerPage,
|
||||
setResourcesSort,
|
||||
resetResourcesFilters,
|
||||
} = resourcesSlice.actions;
|
||||
|
||||
export const resourcesReducer = resourcesSlice.reducer;
|
||||
@@ -0,0 +1,409 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import type { BaseQueryFn } from '@reduxjs/toolkit/query';
|
||||
import { createApi, fetchBaseQuery, type FetchArgs, type FetchBaseQueryError } from '@reduxjs/toolkit/query/react';
|
||||
import type { RootState } from '@/store';
|
||||
import { camelCaseKeys } from '@/utils/format';
|
||||
import type {
|
||||
Attempt,
|
||||
PaginatedResponse,
|
||||
Resources,
|
||||
Rollout,
|
||||
RolloutMode,
|
||||
RolloutStatus,
|
||||
Span,
|
||||
Timestamp,
|
||||
Worker,
|
||||
WorkerStatus,
|
||||
} from '../../types';
|
||||
|
||||
const rawBaseQuery = fetchBaseQuery({
|
||||
baseUrl: '/',
|
||||
});
|
||||
|
||||
const buildAbsoluteUrl = (baseUrl: string, path: string) => {
|
||||
if (path.startsWith('http://') || path.startsWith('https://')) {
|
||||
return path;
|
||||
}
|
||||
|
||||
const normalizedBase = baseUrl.replace(/\/+$/, '');
|
||||
const normalizedPath = path.replace(/^\/+/, '');
|
||||
if (!normalizedBase) {
|
||||
return `/${normalizedPath}`;
|
||||
}
|
||||
return `${normalizedBase}/${normalizedPath}`;
|
||||
};
|
||||
|
||||
const normalizeHeartbeat = (
|
||||
attempt: Partial<Attempt> & { lastHeartbeatTime?: Timestamp | null; lastHeartBeatTime?: Timestamp | null },
|
||||
): Timestamp | null => {
|
||||
if (typeof attempt.lastHeartbeatTime === 'number') {
|
||||
return attempt.lastHeartbeatTime;
|
||||
}
|
||||
|
||||
if (typeof attempt.lastHeartBeatTime === 'number') {
|
||||
return attempt.lastHeartBeatTime;
|
||||
}
|
||||
|
||||
if (typeof attempt.startTime === 'number') {
|
||||
return attempt.startTime;
|
||||
}
|
||||
|
||||
return null;
|
||||
};
|
||||
|
||||
const normalizeAttempt = (value: unknown): Attempt | null => {
|
||||
if (value === null || typeof value === 'undefined') {
|
||||
return null;
|
||||
}
|
||||
|
||||
const camelized = camelCaseKeys(value) as Attempt & {
|
||||
lastHeartbeatTime?: Timestamp | null;
|
||||
lastHeartBeatTime?: Timestamp | null;
|
||||
};
|
||||
const { lastHeartbeatTime, lastHeartBeatTime, ...rest } = camelized;
|
||||
|
||||
return {
|
||||
...rest,
|
||||
lastHeartbeatTime: normalizeHeartbeat({ ...rest, lastHeartbeatTime, lastHeartBeatTime }),
|
||||
};
|
||||
};
|
||||
|
||||
const normalizeAttemptStrict = (value: unknown): Attempt => {
|
||||
const normalized = normalizeAttempt(value);
|
||||
if (!normalized) {
|
||||
throw new Error('Expected attempt payload');
|
||||
}
|
||||
return normalized;
|
||||
};
|
||||
|
||||
const normalizeRollout = (value: unknown): Rollout => {
|
||||
const camelized = camelCaseKeys(value) as Rollout & { attempt?: unknown };
|
||||
const { attempt, ...rest } = camelized;
|
||||
|
||||
return {
|
||||
...rest,
|
||||
attempt: normalizeAttempt(attempt),
|
||||
};
|
||||
};
|
||||
|
||||
const normalizeSpan = (value: unknown): Span => {
|
||||
const camelized = camelCaseKeys(value) as Span & {
|
||||
status?: {
|
||||
status_code?: Span['status']['status_code'];
|
||||
statusCode?: Span['status']['status_code'];
|
||||
description?: string | null;
|
||||
};
|
||||
};
|
||||
const rawStatus = camelized.status ?? { status_code: 'UNSET', description: null };
|
||||
const result = {
|
||||
...camelized,
|
||||
parentId: camelized.parentId ?? null,
|
||||
// The following fields does not need to be normalized to camel case
|
||||
// For example, gen_ai.xxx should not become genAi.xxx
|
||||
attributes: (value as any).attributes ?? {},
|
||||
context: (value as any).context ?? {},
|
||||
parent: (value as any).parent ?? null,
|
||||
resource: (value as any).resource ?? {},
|
||||
status: {
|
||||
status_code: rawStatus.status_code ?? rawStatus.statusCode ?? 'UNSET',
|
||||
description: rawStatus.description ?? null,
|
||||
},
|
||||
};
|
||||
return result;
|
||||
};
|
||||
|
||||
const normalizeResources = (value: unknown): Resources => {
|
||||
const camelized = camelCaseKeys(value) as Resources;
|
||||
return {
|
||||
resourcesId: camelized.resourcesId,
|
||||
version: camelized.version,
|
||||
createTime: camelized.createTime,
|
||||
updateTime: camelized.updateTime,
|
||||
resources: camelized.resources ?? {},
|
||||
};
|
||||
};
|
||||
|
||||
const normalizeWorker = (value: unknown): Worker => {
|
||||
const camelized = camelCaseKeys(value) as Worker;
|
||||
return {
|
||||
workerId: camelized.workerId,
|
||||
status: camelized.status,
|
||||
heartbeatStats: camelized.heartbeatStats ?? null,
|
||||
lastHeartbeatTime: camelized.lastHeartbeatTime ?? null,
|
||||
lastDequeueTime: camelized.lastDequeueTime ?? null,
|
||||
lastBusyTime: camelized.lastBusyTime ?? null,
|
||||
lastIdleTime: camelized.lastIdleTime ?? null,
|
||||
currentRolloutId: camelized.currentRolloutId ?? null,
|
||||
currentAttemptId: camelized.currentAttemptId ?? null,
|
||||
};
|
||||
};
|
||||
|
||||
const normalizePaginatedResponse = <T>(value: unknown, normalizer: (item: unknown) => T): PaginatedResponse<T> => {
|
||||
if (!value || typeof value !== 'object') {
|
||||
throw new Error('Expected paginated response payload');
|
||||
}
|
||||
|
||||
const converted = value as {
|
||||
items?: unknown;
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
total?: number;
|
||||
};
|
||||
|
||||
const itemsSource = Array.isArray(converted.items) ? converted.items : [];
|
||||
|
||||
return {
|
||||
items: itemsSource.map((item) => normalizer(item)),
|
||||
limit: typeof converted.limit === 'number' ? converted.limit : itemsSource.length,
|
||||
offset: typeof converted.offset === 'number' ? converted.offset : 0,
|
||||
total: typeof converted.total === 'number' ? converted.total : itemsSource.length,
|
||||
};
|
||||
};
|
||||
|
||||
const dynamicBaseQuery: BaseQueryFn<string | FetchArgs, unknown, FetchBaseQueryError> = async (
|
||||
args,
|
||||
api,
|
||||
extraOptions,
|
||||
) => {
|
||||
const state = api.getState() as RootState;
|
||||
const stateBaseUrl = state.config?.baseUrl;
|
||||
const fallbackBaseUrl = typeof window !== 'undefined' ? window.location.origin : '';
|
||||
const baseUrl = stateBaseUrl && stateBaseUrl.trim().length > 0 ? stateBaseUrl : fallbackBaseUrl;
|
||||
const preparedArgs: FetchArgs =
|
||||
typeof args === 'string'
|
||||
? { url: args }
|
||||
: {
|
||||
...args,
|
||||
url: args.url ?? '',
|
||||
};
|
||||
|
||||
const absoluteUrl = buildAbsoluteUrl(baseUrl, preparedArgs.url ?? '');
|
||||
return rawBaseQuery({ ...preparedArgs, url: absoluteUrl }, api, extraOptions);
|
||||
};
|
||||
|
||||
export type GetRolloutsQueryArgs = {
|
||||
limit: number;
|
||||
offset: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
statusIn?: RolloutStatus[];
|
||||
rolloutIdContains?: string | null;
|
||||
modeIn?: RolloutMode[];
|
||||
};
|
||||
|
||||
export type GetResourcesQueryArgs = {
|
||||
limit: number;
|
||||
offset: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
resourcesIdContains?: string | null;
|
||||
};
|
||||
|
||||
export type GetWorkersQueryArgs = {
|
||||
limit: number;
|
||||
offset: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
workerIdContains?: string | null;
|
||||
statusIn?: WorkerStatus[];
|
||||
};
|
||||
|
||||
export type GetRolloutAttemptsQueryArgs = {
|
||||
rolloutId: string;
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
};
|
||||
|
||||
export type GetSpansQueryArgs = {
|
||||
rolloutId: string;
|
||||
attemptId?: string | null;
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
sortBy?: string | null;
|
||||
sortOrder?: 'asc' | 'desc';
|
||||
traceIdContains?: string | null;
|
||||
spanIdContains?: string | null;
|
||||
parentIdContains?: string | null;
|
||||
nameContains?: string | null;
|
||||
filterLogic?: 'and' | 'or' | null;
|
||||
};
|
||||
|
||||
export const rolloutsApi = createApi({
|
||||
reducerPath: 'rolloutsApi',
|
||||
baseQuery: dynamicBaseQuery,
|
||||
tagTypes: ['Rollout', 'Span', 'Resources', 'Worker'],
|
||||
endpoints: (builder) => ({
|
||||
getResources: builder.query<PaginatedResponse<Resources>, GetResourcesQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, resourcesIdContains }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
searchParams.set('limit', String(typeof limit === 'number' ? limit : -1));
|
||||
searchParams.set('offset', String(typeof offset === 'number' ? offset : 0));
|
||||
if (sortBy) {
|
||||
searchParams.set('sort_by', sortBy);
|
||||
}
|
||||
if (sortOrder) {
|
||||
searchParams.set('sort_order', sortOrder);
|
||||
}
|
||||
if (resourcesIdContains && resourcesIdContains.trim().length > 0) {
|
||||
searchParams.set('resources_id_contains', resourcesIdContains.trim());
|
||||
}
|
||||
|
||||
const queryString = searchParams.toString();
|
||||
const url = queryString.length > 0 ? `v1/agl/resources?${queryString}` : 'v1/agl/resources';
|
||||
return { url, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeResources),
|
||||
providesTags: (result) =>
|
||||
result
|
||||
? [
|
||||
{ type: 'Resources' as const, id: 'LIST' },
|
||||
...result.items.map((item) => ({ type: 'Resources' as const, id: item.resourcesId })),
|
||||
]
|
||||
: [{ type: 'Resources' as const, id: 'LIST' }],
|
||||
}),
|
||||
getWorkers: builder.query<PaginatedResponse<Worker>, GetWorkersQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, workerIdContains, statusIn }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
searchParams.set('limit', String(typeof limit === 'number' ? limit : -1));
|
||||
searchParams.set('offset', String(typeof offset === 'number' ? offset : 0));
|
||||
if (sortBy) {
|
||||
searchParams.set('sort_by', sortBy);
|
||||
}
|
||||
if (sortOrder) {
|
||||
searchParams.set('sort_order', sortOrder);
|
||||
}
|
||||
if (workerIdContains && workerIdContains.trim().length > 0) {
|
||||
searchParams.set('worker_id_contains', workerIdContains.trim());
|
||||
}
|
||||
if (statusIn && statusIn.length > 0) {
|
||||
statusIn.forEach((status) => searchParams.append('status_in', status));
|
||||
}
|
||||
|
||||
const queryString = searchParams.toString();
|
||||
const url = queryString.length > 0 ? `v1/agl/workers?${queryString}` : 'v1/agl/workers';
|
||||
return { url, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeWorker),
|
||||
providesTags: (result) =>
|
||||
result
|
||||
? [
|
||||
{ type: 'Worker' as const, id: 'LIST' },
|
||||
...result.items.map((worker) => ({ type: 'Worker' as const, id: worker.workerId })),
|
||||
]
|
||||
: [{ type: 'Worker' as const, id: 'LIST' }],
|
||||
}),
|
||||
getRollouts: builder.query<PaginatedResponse<Rollout>, GetRolloutsQueryArgs>({
|
||||
query: ({ limit, offset, sortBy, sortOrder, statusIn, rolloutIdContains, modeIn }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
searchParams.set('limit', String(typeof limit === 'number' ? limit : -1));
|
||||
searchParams.set('offset', String(typeof offset === 'number' ? offset : 0));
|
||||
if (sortBy) {
|
||||
searchParams.set('sort_by', sortBy);
|
||||
}
|
||||
if (sortOrder) {
|
||||
searchParams.set('sort_order', sortOrder);
|
||||
}
|
||||
if (statusIn && statusIn.length > 0) {
|
||||
statusIn.forEach((status) => searchParams.append('status_in', status));
|
||||
}
|
||||
if (modeIn && modeIn.length > 0) {
|
||||
modeIn.forEach((mode) => searchParams.append('mode_in', mode));
|
||||
}
|
||||
if (rolloutIdContains && rolloutIdContains.trim().length > 0) {
|
||||
searchParams.set('rollout_id_contains', rolloutIdContains.trim());
|
||||
}
|
||||
|
||||
const queryString = searchParams.toString();
|
||||
const url = queryString.length > 0 ? `v1/agl/rollouts?${queryString}` : 'v1/agl/rollouts';
|
||||
return { url, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeRollout),
|
||||
providesTags: (result) =>
|
||||
result
|
||||
? [
|
||||
{ type: 'Rollout' as const, id: 'LIST' },
|
||||
...result.items.map((rollout) => ({ type: 'Rollout' as const, id: rollout.rolloutId })),
|
||||
]
|
||||
: [{ type: 'Rollout' as const, id: 'LIST' }],
|
||||
}),
|
||||
getRolloutAttempts: builder.query<PaginatedResponse<Attempt>, GetRolloutAttemptsQueryArgs>({
|
||||
query: ({ rolloutId, limit = -1, offset = 0, sortBy, sortOrder }) => {
|
||||
const searchParams = new URLSearchParams();
|
||||
searchParams.set('limit', String(typeof limit === 'number' ? limit : -1));
|
||||
searchParams.set('offset', String(typeof offset === 'number' ? offset : 0));
|
||||
if (sortBy) {
|
||||
searchParams.set('sort_by', sortBy);
|
||||
}
|
||||
if (sortOrder) {
|
||||
searchParams.set('sort_order', sortOrder);
|
||||
}
|
||||
const queryString = searchParams.toString();
|
||||
const url =
|
||||
queryString.length > 0
|
||||
? `v1/agl/rollouts/${rolloutId}/attempts?${queryString}`
|
||||
: `v1/agl/rollouts/${rolloutId}/attempts`;
|
||||
return { url, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeAttemptStrict),
|
||||
providesTags: (_result, _error, queryArgs) => [{ type: 'Rollout', id: queryArgs.rolloutId }],
|
||||
}),
|
||||
getSpans: builder.query<PaginatedResponse<Span>, GetSpansQueryArgs>({
|
||||
query: (args) => {
|
||||
if (!args.rolloutId) {
|
||||
throw new Error('rolloutId is required to fetch spans');
|
||||
}
|
||||
const searchParams = new URLSearchParams({ rollout_id: args.rolloutId });
|
||||
if (args.attemptId) {
|
||||
searchParams.set('attempt_id', args.attemptId);
|
||||
}
|
||||
if (typeof args.limit === 'number') {
|
||||
searchParams.set('limit', String(args.limit));
|
||||
}
|
||||
if (typeof args.offset === 'number') {
|
||||
searchParams.set('offset', String(args.offset));
|
||||
}
|
||||
if (args.sortBy) {
|
||||
searchParams.set('sort_by', args.sortBy);
|
||||
}
|
||||
if (args.sortOrder) {
|
||||
searchParams.set('sort_order', args.sortOrder);
|
||||
}
|
||||
if (args.traceIdContains) {
|
||||
searchParams.set('trace_id_contains', args.traceIdContains);
|
||||
}
|
||||
if (args.spanIdContains) {
|
||||
searchParams.set('span_id_contains', args.spanIdContains);
|
||||
}
|
||||
if (args.parentIdContains) {
|
||||
searchParams.set('parent_id_contains', args.parentIdContains);
|
||||
}
|
||||
if (args.nameContains) {
|
||||
searchParams.set('name_contains', args.nameContains);
|
||||
}
|
||||
if (args.filterLogic) {
|
||||
searchParams.set('filter_logic', args.filterLogic);
|
||||
}
|
||||
return { url: `v1/agl/spans?${searchParams.toString()}`, method: 'GET' };
|
||||
},
|
||||
transformResponse: (response: unknown) => normalizePaginatedResponse(response, normalizeSpan),
|
||||
providesTags: (_result, _error, args) =>
|
||||
args
|
||||
? [
|
||||
{ type: 'Span' as const, id: `${args.rolloutId}:${args.attemptId ?? 'latest'}` },
|
||||
{ type: 'Span' as const, id: 'LIST' },
|
||||
]
|
||||
: [{ type: 'Span' as const, id: 'LIST' }],
|
||||
}),
|
||||
}),
|
||||
});
|
||||
|
||||
export const {
|
||||
useGetResourcesQuery,
|
||||
useGetWorkersQuery,
|
||||
useGetRolloutsQuery,
|
||||
useGetRolloutAttemptsQuery,
|
||||
useGetSpansQuery,
|
||||
} = rolloutsApi;
|
||||
@@ -0,0 +1,6 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
export * from './api';
|
||||
export * from './slice';
|
||||
export * from './selectors';
|
||||
export * from '../../types';
|
||||
@@ -0,0 +1,127 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import { createServerBackedStore } from '@test-utils';
|
||||
import { describe, expect, it } from 'vitest';
|
||||
import { rolloutsApi } from './api';
|
||||
import { selectRolloutsQueryArgs } from './selectors';
|
||||
import {
|
||||
resetRolloutsFilters,
|
||||
setRolloutsModeFilters,
|
||||
setRolloutsPage,
|
||||
setRolloutsRecordsPerPage,
|
||||
setRolloutsSearchTerm,
|
||||
setRolloutsSort,
|
||||
setRolloutsStatusFilters,
|
||||
} from './slice';
|
||||
|
||||
describe('rollouts feature integration', () => {
|
||||
it('builds default query arguments from the UI state', () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectRolloutsQueryArgs(store.getState());
|
||||
|
||||
expect(queryArgs).toMatchObject({
|
||||
limit: 100,
|
||||
offset: 0,
|
||||
sortBy: 'start_time',
|
||||
sortOrder: 'desc',
|
||||
rolloutIdContains: undefined,
|
||||
statusIn: undefined,
|
||||
modeIn: undefined,
|
||||
});
|
||||
});
|
||||
|
||||
it('retrieves rollouts from the Python LightningStore server', async () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectRolloutsQueryArgs(store.getState());
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getRollouts.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.total).toBe(6);
|
||||
expect(data.items).toHaveLength(6);
|
||||
|
||||
const rolloutIds = data.items.map((rollout) => rollout.rolloutId);
|
||||
expect(rolloutIds).toEqual(
|
||||
expect.arrayContaining(['ro-story-001', 'ro-story-002', 'ro-story-003', 'ro-story-004', 'ro-story-005']),
|
||||
);
|
||||
|
||||
const startTimes = data.items.map((rollout) => rollout.startTime);
|
||||
const sortedStartTimes = [...startTimes].sort((a, b) => b - a);
|
||||
expect(startTimes).toEqual(sortedStartTimes);
|
||||
expect(data.items[0].rolloutId).toBe('ro-story-005');
|
||||
expect(data.items[0].status).toBeDefined();
|
||||
});
|
||||
|
||||
it('includes attempts directly on rollout payloads when they exist', async () => {
|
||||
const store = createServerBackedStore();
|
||||
const queryArgs = selectRolloutsQueryArgs(store.getState());
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getRollouts.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
const rolloutWithAttempt = data.items.find((rollout) => rollout.rolloutId === 'ro-story-002');
|
||||
expect(rolloutWithAttempt).toBeDefined();
|
||||
expect(rolloutWithAttempt?.attempt).not.toBeNull();
|
||||
expect(rolloutWithAttempt?.attempt?.attemptId).toBe('at-story-022');
|
||||
|
||||
const rolloutWithoutAttempt = data.items.find((rollout) => rollout.rolloutId === 'ro-story-004');
|
||||
expect(rolloutWithoutAttempt).toBeDefined();
|
||||
expect(rolloutWithoutAttempt?.attempt).toBeNull();
|
||||
});
|
||||
|
||||
it('retrieves attempts for a rollout from the Python server', async () => {
|
||||
const store = createServerBackedStore();
|
||||
const subscription = store.dispatch(
|
||||
rolloutsApi.endpoints.getRolloutAttempts.initiate({ rolloutId: 'ro-story-002' }),
|
||||
);
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.total).toBe(2);
|
||||
expect(data.items.map((attempt) => attempt.attemptId)).toEqual(['at-story-021', 'at-story-022']);
|
||||
});
|
||||
|
||||
it('paginates rollouts with custom UI state', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(setRolloutsRecordsPerPage(2));
|
||||
store.dispatch(setRolloutsPage(2));
|
||||
|
||||
const queryArgs = selectRolloutsQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({ limit: 2, offset: 2 });
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getRollouts.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(2);
|
||||
expect(data.items.map((rollout) => rollout.rolloutId)).toEqual(['ro-story-004', 'ro-story-002']);
|
||||
});
|
||||
|
||||
it('filters and sorts rollouts based on UI selections', async () => {
|
||||
const store = createServerBackedStore();
|
||||
store.dispatch(resetRolloutsFilters());
|
||||
store.dispatch(setRolloutsStatusFilters(['succeeded']));
|
||||
store.dispatch(setRolloutsModeFilters(['val']));
|
||||
store.dispatch(setRolloutsSearchTerm('ro-story-002'));
|
||||
store.dispatch(setRolloutsSort({ column: 'rolloutId', direction: 'asc' }));
|
||||
|
||||
const queryArgs = selectRolloutsQueryArgs(store.getState());
|
||||
expect(queryArgs).toMatchObject({
|
||||
statusIn: ['succeeded'],
|
||||
modeIn: ['val'],
|
||||
rolloutIdContains: 'ro-story-002',
|
||||
sortBy: 'rollout_id',
|
||||
sortOrder: 'asc',
|
||||
});
|
||||
|
||||
const subscription = store.dispatch(rolloutsApi.endpoints.getRollouts.initiate(queryArgs));
|
||||
const data = await subscription.unwrap();
|
||||
subscription.unsubscribe();
|
||||
|
||||
expect(data.items).toHaveLength(1);
|
||||
expect(data.items[0].rolloutId).toBe('ro-story-002');
|
||||
expect(data.items[0].status).toBe('succeeded');
|
||||
});
|
||||
});
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user