nat-engine 1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mannf/__init__.py +33 -0
- mannf/__main__.py +10 -0
- mannf/_version.py +8 -0
- mannf/agents/__init__.py +7 -0
- mannf/agents/analyzer_agent.py +9 -0
- mannf/agents/base.py +9 -0
- mannf/agents/bdi_agent.py +9 -0
- mannf/agents/belief_state.py +9 -0
- mannf/agents/coordinator_agent.py +9 -0
- mannf/agents/executor_agent.py +9 -0
- mannf/agents/monitor_agent.py +9 -0
- mannf/agents/oracle_agent.py +9 -0
- mannf/agents/planner_agent.py +9 -0
- mannf/agents/test_agent.py +9 -0
- mannf/anomaly/__init__.py +7 -0
- mannf/anomaly/enhanced_detector.py +9 -0
- mannf/cli.py +9 -0
- mannf/core/__init__.py +26 -0
- mannf/core/agents/__init__.py +52 -0
- mannf/core/agents/accessibility_scanner_agent.py +245 -0
- mannf/core/agents/analyzer_agent.py +224 -0
- mannf/core/agents/autonomous_loop_agent.py +1086 -0
- mannf/core/agents/autonomous_loop_models.py +62 -0
- mannf/core/agents/autonomous_run_differ.py +427 -0
- mannf/core/agents/base.py +128 -0
- mannf/core/agents/bdi_agent.py +330 -0
- mannf/core/agents/belief_state.py +202 -0
- mannf/core/agents/browser_coordinator_agent.py +224 -0
- mannf/core/agents/browser_executor_agent.py +410 -0
- mannf/core/agents/coordinator_agent.py +262 -0
- mannf/core/agents/executor_agent.py +222 -0
- mannf/core/agents/monitor_agent.py +188 -0
- mannf/core/agents/oracle_agent.py +150 -0
- mannf/core/agents/performance_testing_agent.py +279 -0
- mannf/core/agents/planner_agent.py +128 -0
- mannf/core/agents/test_agent.py +249 -0
- mannf/core/agents/visual_regression_agent.py +311 -0
- mannf/core/agents/web_crawler_agent.py +510 -0
- mannf/core/agents/worker_pool.py +366 -0
- mannf/core/anomaly/__init__.py +14 -0
- mannf/core/anomaly/enhanced_detector.py +541 -0
- mannf/core/browser/__init__.py +63 -0
- mannf/core/browser/accessibility_scanner.py +424 -0
- mannf/core/browser/discovery_model.py +178 -0
- mannf/core/browser/dom_snapshot.py +349 -0
- mannf/core/browser/ingestor_bridge.py +371 -0
- mannf/core/browser/performance_metrics.py +217 -0
- mannf/core/browser/reflection_analyzer.py +442 -0
- mannf/core/browser/scenario_generator.py +1100 -0
- mannf/core/browser/security_scenario_generator.py +695 -0
- mannf/core/browser/visual_comparer.py +159 -0
- mannf/core/diagnostics/__init__.py +28 -0
- mannf/core/diagnostics/failure_clusterer.py +211 -0
- mannf/core/diagnostics/flake_detector.py +233 -0
- mannf/core/diagnostics/root_cause_analyzer.py +273 -0
- mannf/core/distributed/__init__.py +16 -0
- mannf/core/distributed/endpoint.py +139 -0
- mannf/core/distributed/system_under_test.py +207 -0
- mannf/core/functional_orchestrator.py +428 -0
- mannf/core/messaging/__init__.py +11 -0
- mannf/core/messaging/bus.py +113 -0
- mannf/core/messaging/messages.py +89 -0
- mannf/core/nat_orchestrator.py +342 -0
- mannf/core/neural/__init__.py +183 -0
- mannf/core/orchestrator.py +272 -0
- mannf/core/prioritization/__init__.py +17 -0
- mannf/core/prioritization/adaptive_controller.py +509 -0
- mannf/core/prioritization/belief_prioritizer.py +231 -0
- mannf/core/prioritization/risk_scorer.py +430 -0
- mannf/core/reporting/__init__.py +12 -0
- mannf/core/reporting/unified_report.py +664 -0
- mannf/core/testing/__init__.py +17 -0
- mannf/core/testing/adaptive_controller.py +149 -0
- mannf/core/testing/models.py +179 -0
- mannf/core/validation/__init__.py +10 -0
- mannf/core/validation/self_validation_runner.py +180 -0
- mannf/dashboard/__init__.py +7 -0
- mannf/dashboard/app.py +9 -0
- mannf/dashboard/models.py +9 -0
- mannf/dashboard/static/index.html +2538 -0
- mannf/dashboard/telemetry.py +9 -0
- mannf/distributed/__init__.py +7 -0
- mannf/distributed/endpoint.py +9 -0
- mannf/distributed/system_under_test.py +9 -0
- mannf/healing/__init__.py +7 -0
- mannf/healing/graphql_schema_diff.py +9 -0
- mannf/healing/healer.py +9 -0
- mannf/healing/models.py +9 -0
- mannf/healing/schema_diff.py +9 -0
- mannf/integrations/__init__.py +7 -0
- mannf/integrations/auth.py +9 -0
- mannf/integrations/graphql_parser.py +9 -0
- mannf/integrations/graphql_sut.py +9 -0
- mannf/integrations/http_sut.py +9 -0
- mannf/integrations/openapi_parser.py +9 -0
- mannf/integrations/postman_parser.py +9 -0
- mannf/llm/__init__.py +7 -0
- mannf/llm/anthropic_provider.py +9 -0
- mannf/llm/base.py +9 -0
- mannf/llm/config.py +9 -0
- mannf/llm/factory.py +9 -0
- mannf/llm/openai_provider.py +9 -0
- mannf/llm/prompts.py +9 -0
- mannf/messaging/__init__.py +7 -0
- mannf/messaging/bus.py +9 -0
- mannf/messaging/messages.py +9 -0
- mannf/nat_orchestrator.py +9 -0
- mannf/neural/__init__.py +7 -0
- mannf/orchestrator.py +9 -0
- mannf/prioritization/__init__.py +7 -0
- mannf/prioritization/adaptive_controller.py +9 -0
- mannf/prioritization/belief_prioritizer.py +9 -0
- mannf/prioritization/risk_scorer.py +9 -0
- mannf/product/__init__.py +29 -0
- mannf/product/admin/__init__.py +3 -0
- mannf/product/admin/routes.py +514 -0
- mannf/product/auth/__init__.py +5 -0
- mannf/product/auth/saml.py +212 -0
- mannf/product/billing/__init__.py +5 -0
- mannf/product/billing/audit.py +160 -0
- mannf/product/billing/feature_gates.py +180 -0
- mannf/product/billing/metering.py +179 -0
- mannf/product/billing/notifications.py +181 -0
- mannf/product/billing/plans.py +133 -0
- mannf/product/billing/rate_limits.py +35 -0
- mannf/product/billing/stripe_billing.py +906 -0
- mannf/product/billing/tenant_auth.py +233 -0
- mannf/product/billing/tenant_manager.py +873 -0
- mannf/product/cli.py +3900 -0
- mannf/product/cli_admin.py +408 -0
- mannf/product/dashboard/__init__.py +61 -0
- mannf/product/dashboard/app.py +3567 -0
- mannf/product/dashboard/models.py +460 -0
- mannf/product/dashboard/static/index.html +6347 -0
- mannf/product/dashboard/static/manifest.json +25 -0
- mannf/product/dashboard/static/pwa-icon-192.png +0 -0
- mannf/product/dashboard/static/pwa-icon-512.png +0 -0
- mannf/product/dashboard/static/sw.js +64 -0
- mannf/product/dashboard/telemetry.py +547 -0
- mannf/product/database.py +145 -0
- mannf/product/demo.py +844 -0
- mannf/product/doctor.py +509 -0
- mannf/product/exporters/__init__.py +65 -0
- mannf/product/exporters/azuredevops_exporter.py +257 -0
- mannf/product/exporters/base.py +307 -0
- mannf/product/exporters/bugzilla_exporter.py +200 -0
- mannf/product/exporters/dedup.py +275 -0
- mannf/product/exporters/finding_adapter.py +216 -0
- mannf/product/exporters/github_exporter.py +197 -0
- mannf/product/exporters/gitlab_exporter.py +215 -0
- mannf/product/exporters/jira_exporter.py +180 -0
- mannf/product/exporters/linear_exporter.py +195 -0
- mannf/product/exporters/loader.py +233 -0
- mannf/product/exporters/pagerduty_exporter.py +363 -0
- mannf/product/exporters/sentry_exporter.py +322 -0
- mannf/product/exporters/servicenow_exporter.py +240 -0
- mannf/product/exporters/shortcut_exporter.py +231 -0
- mannf/product/exporters/webhook_exporter.py +383 -0
- mannf/product/formatters/__init__.py +18 -0
- mannf/product/formatters/allure_formatter.py +161 -0
- mannf/product/formatters/ctrf_formatter.py +149 -0
- mannf/product/healing/__init__.py +30 -0
- mannf/product/healing/graphql_schema_diff.py +152 -0
- mannf/product/healing/healer.py +141 -0
- mannf/product/healing/models.py +175 -0
- mannf/product/healing/schema_diff.py +251 -0
- mannf/product/ingestors/__init__.py +77 -0
- mannf/product/ingestors/base.py +256 -0
- mannf/product/ingestors/bgstm_ingestor.py +764 -0
- mannf/product/ingestors/curl_ingestor.py +1019 -0
- mannf/product/ingestors/cypress_ingestor.py +487 -0
- mannf/product/ingestors/gherkin_ingestor.py +967 -0
- mannf/product/ingestors/graphql_ingestor.py +845 -0
- mannf/product/ingestors/grpc_ingestor.py +591 -0
- mannf/product/ingestors/har_ingestor.py +976 -0
- mannf/product/ingestors/loader.py +284 -0
- mannf/product/ingestors/models.py +146 -0
- mannf/product/ingestors/openapi_ingestor.py +606 -0
- mannf/product/ingestors/playwright_ingestor.py +449 -0
- mannf/product/ingestors/postman_ingestor.py +631 -0
- mannf/product/ingestors/traffic_ingestor.py +679 -0
- mannf/product/ingestors/websocket_ingestor.py +526 -0
- mannf/product/integrations/__init__.py +21 -0
- mannf/product/integrations/auth.py +190 -0
- mannf/product/integrations/graphql_parser.py +436 -0
- mannf/product/integrations/graphql_sut.py +247 -0
- mannf/product/integrations/grpc_sut.py +469 -0
- mannf/product/integrations/http_sut.py +237 -0
- mannf/product/integrations/kafka_adapter.py +342 -0
- mannf/product/integrations/openapi_parser.py +513 -0
- mannf/product/integrations/postman_parser.py +467 -0
- mannf/product/integrations/webhook_receiver.py +344 -0
- mannf/product/integrations/websocket_sut.py +434 -0
- mannf/product/llm/__init__.py +25 -0
- mannf/product/llm/anthropic_provider.py +94 -0
- mannf/product/llm/base.py +267 -0
- mannf/product/llm/config.py +48 -0
- mannf/product/llm/factory.py +42 -0
- mannf/product/llm/openai_provider.py +93 -0
- mannf/product/llm/prompts.py +403 -0
- mannf/product/llm/root_cause_service.py +311 -0
- mannf/product/llm/test_plan_models.py +78 -0
- mannf/product/metrics.py +149 -0
- mannf/product/middleware/__init__.py +3 -0
- mannf/product/middleware/audit_middleware.py +112 -0
- mannf/product/middleware/tenant_isolation.py +114 -0
- mannf/product/models.py +347 -0
- mannf/product/notifications/__init__.py +24 -0
- mannf/product/notifications/dispatcher.py +411 -0
- mannf/product/onboarding.py +190 -0
- mannf/product/orchestration/__init__.py +39 -0
- mannf/product/orchestration/ingest_scan_orchestrator.py +339 -0
- mannf/product/orchestration/pipeline.py +401 -0
- mannf/product/orchestrator.py +987 -0
- mannf/product/orchestrator_models.py +269 -0
- mannf/product/regression/__init__.py +36 -0
- mannf/product/regression/differ.py +172 -0
- mannf/product/regression/masking.py +100 -0
- mannf/product/regression/models.py +232 -0
- mannf/product/regression/recorder.py +124 -0
- mannf/product/regression/replayer.py +168 -0
- mannf/product/reports/__init__.py +10 -0
- mannf/product/reports/pdf.py +132 -0
- mannf/product/scheduling/__init__.py +57 -0
- mannf/product/scheduling/cron_utils.py +251 -0
- mannf/product/scheduling/engine.py +473 -0
- mannf/product/scheduling/models.py +86 -0
- mannf/product/scheduling/queue.py +894 -0
- mannf/product/scheduling/store.py +235 -0
- mannf/product/security/__init__.py +21 -0
- mannf/product/security/belief_guided.py +143 -0
- mannf/product/security/checks/__init__.py +55 -0
- mannf/product/security/checks/base.py +69 -0
- mannf/product/security/checks/bfla.py +77 -0
- mannf/product/security/checks/bola.py +77 -0
- mannf/product/security/checks/bopla.py +80 -0
- mannf/product/security/checks/broken_auth.py +86 -0
- mannf/product/security/checks/graphql_security.py +299 -0
- mannf/product/security/checks/inventory.py +70 -0
- mannf/product/security/checks/misconfig.py +158 -0
- mannf/product/security/checks/resource_consumption.py +70 -0
- mannf/product/security/checks/sensitive_flows.py +80 -0
- mannf/product/security/checks/ssrf.py +101 -0
- mannf/product/security/checks/unsafe_consumption.py +120 -0
- mannf/product/security/models.py +92 -0
- mannf/product/security/plugin_loader.py +182 -0
- mannf/product/security/reporter.py +92 -0
- mannf/product/security/scanner.py +183 -0
- mannf/product/server.py +6220 -0
- mannf/product/setup_wizard.py +873 -0
- mannf/product/status.py +404 -0
- mannf/product/storage/__init__.py +10 -0
- mannf/product/storage/artifact_store.py +343 -0
- mannf/product/telemetry.py +300 -0
- mannf/product/uninstall.py +169 -0
- mannf/product/upgrade.py +139 -0
- mannf/product/weights/__init__.py +13 -0
- mannf/product/weights/blob_store.py +299 -0
- mannf/product/weights/factory.py +42 -0
- mannf/product/weights/registry.py +159 -0
- mannf/product/weights/store.py +210 -0
- mannf/regression/__init__.py +7 -0
- mannf/regression/differ.py +9 -0
- mannf/regression/masking.py +9 -0
- mannf/regression/models.py +9 -0
- mannf/regression/recorder.py +9 -0
- mannf/regression/replayer.py +9 -0
- mannf/security/__init__.py +7 -0
- mannf/security/belief_guided.py +9 -0
- mannf/security/checks/__init__.py +7 -0
- mannf/security/checks/base.py +9 -0
- mannf/security/checks/bfla.py +9 -0
- mannf/security/checks/bola.py +9 -0
- mannf/security/checks/bopla.py +9 -0
- mannf/security/checks/broken_auth.py +9 -0
- mannf/security/checks/graphql_security.py +9 -0
- mannf/security/checks/inventory.py +9 -0
- mannf/security/checks/misconfig.py +9 -0
- mannf/security/checks/resource_consumption.py +9 -0
- mannf/security/checks/sensitive_flows.py +9 -0
- mannf/security/checks/ssrf.py +9 -0
- mannf/security/checks/unsafe_consumption.py +9 -0
- mannf/security/models.py +9 -0
- mannf/security/reporter.py +9 -0
- mannf/security/scanner.py +9 -0
- mannf/server.py +9 -0
- mannf/testing/__init__.py +7 -0
- mannf/testing/adaptive_controller.py +9 -0
- mannf/testing/models.py +9 -0
- mannf/weights/__init__.py +7 -0
- mannf/weights/registry.py +9 -0
- mannf/weights/store.py +9 -0
- nat_engine-1.dist-info/METADATA +555 -0
- nat_engine-1.dist-info/RECORD +299 -0
- nat_engine-1.dist-info/WHEEL +5 -0
- nat_engine-1.dist-info/entry_points.txt +4 -0
- nat_engine-1.dist-info/licenses/LICENSE +651 -0
- nat_engine-1.dist-info/licenses/NOTICE +178 -0
- nat_engine-1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
# Copyright (C) 2026 Brad Guider
|
|
2
|
+
# This file is part of NAT (Neural Agent Testing Framework).
|
|
3
|
+
# Licensed under the AGPL-3.0. See LICENSE for details.
|
|
4
|
+
# Commercial licensing available — see COMMERCIAL_LICENSE.md.
|
|
5
|
+
|
|
6
|
+
"""Uninstall command for NAT — removes config, data, and shell completions.
|
|
7
|
+
|
|
8
|
+
Entry point::
|
|
9
|
+
|
|
10
|
+
def run_uninstall(args: argparse.Namespace) -> int:
|
|
11
|
+
...
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import argparse
|
|
17
|
+
import os
|
|
18
|
+
import shutil
|
|
19
|
+
import subprocess
|
|
20
|
+
import sys
|
|
21
|
+
from pathlib import Path
|
|
22
|
+
|
|
23
|
+
# ---------------------------------------------------------------------------
|
|
24
|
+
# Exit codes
|
|
25
|
+
# ---------------------------------------------------------------------------
|
|
26
|
+
_EXIT_OK = 0
|
|
27
|
+
_EXIT_ABORTED = 1
|
|
28
|
+
_EXIT_ERROR = 2
|
|
29
|
+
|
|
30
|
+
# ---------------------------------------------------------------------------
|
|
31
|
+
# Shell completion file locations
|
|
32
|
+
# ---------------------------------------------------------------------------
|
|
33
|
+
|
|
34
|
+
_BASH_COMPLETION_PATHS = [
|
|
35
|
+
Path("/etc/bash_completion.d/nat"),
|
|
36
|
+
Path("/usr/local/etc/bash_completion.d/nat"),
|
|
37
|
+
]
|
|
38
|
+
|
|
39
|
+
_ZSH_COMPLETION_PATHS = [
|
|
40
|
+
Path("/usr/local/share/zsh/site-functions/_nat"),
|
|
41
|
+
Path("/usr/share/zsh/site-functions/_nat"),
|
|
42
|
+
]
|
|
43
|
+
|
|
44
|
+
_FISH_COMPLETION_PATHS = [
|
|
45
|
+
Path.home() / ".config" / "fish" / "completions" / "nat.fish",
|
|
46
|
+
Path("/usr/share/fish/completions/nat.fish"),
|
|
47
|
+
]
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _remove_path(path: Path, removed: list[str], errors: list[str]) -> None:
|
|
51
|
+
"""Remove *path* (file or directory) and record the outcome."""
|
|
52
|
+
try:
|
|
53
|
+
if path.is_dir():
|
|
54
|
+
shutil.rmtree(path)
|
|
55
|
+
removed.append(str(path))
|
|
56
|
+
elif path.exists():
|
|
57
|
+
path.unlink()
|
|
58
|
+
removed.append(str(path))
|
|
59
|
+
except OSError as exc:
|
|
60
|
+
errors.append(f"{path}: {exc}")
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def run_uninstall(args: argparse.Namespace) -> int:
|
|
64
|
+
"""Remove NAT configuration files, data directory, and shell completions.
|
|
65
|
+
|
|
66
|
+
Parameters
|
|
67
|
+
----------
|
|
68
|
+
args:
|
|
69
|
+
Parsed CLI arguments. The following attributes are read:
|
|
70
|
+
|
|
71
|
+
- ``keep_config`` (bool) — preserve ``.natrc`` files.
|
|
72
|
+
- ``yes`` (bool) — skip confirmation prompt.
|
|
73
|
+
- ``purge`` (bool) — also run ``alembic downgrade base`` and remove
|
|
74
|
+
Docker volumes when ``DATABASE_URL`` is set.
|
|
75
|
+
|
|
76
|
+
Returns
|
|
77
|
+
-------
|
|
78
|
+
int
|
|
79
|
+
Exit code (0 = success, non-zero = error).
|
|
80
|
+
"""
|
|
81
|
+
keep_config: bool = getattr(args, "keep_config", False)
|
|
82
|
+
skip_confirm: bool = getattr(args, "yes", False)
|
|
83
|
+
purge: bool = getattr(args, "purge", False)
|
|
84
|
+
|
|
85
|
+
print()
|
|
86
|
+
print(" ⚠️ NAT Uninstall")
|
|
87
|
+
print(" " + "─" * 48)
|
|
88
|
+
|
|
89
|
+
# -----------------------------------------------------------------------
|
|
90
|
+
# Confirmation
|
|
91
|
+
# -----------------------------------------------------------------------
|
|
92
|
+
if not skip_confirm:
|
|
93
|
+
try:
|
|
94
|
+
answer = input(" Are you sure you want to uninstall NAT? [y/N] ").strip().lower()
|
|
95
|
+
except (EOFError, KeyboardInterrupt):
|
|
96
|
+
print()
|
|
97
|
+
print(" Aborted.")
|
|
98
|
+
return _EXIT_ABORTED
|
|
99
|
+
if answer not in ("y", "yes"):
|
|
100
|
+
print(" Aborted.")
|
|
101
|
+
return _EXIT_ABORTED
|
|
102
|
+
|
|
103
|
+
removed: list[str] = []
|
|
104
|
+
errors: list[str] = []
|
|
105
|
+
|
|
106
|
+
# -----------------------------------------------------------------------
|
|
107
|
+
# Remove .natrc files
|
|
108
|
+
# -----------------------------------------------------------------------
|
|
109
|
+
if not keep_config:
|
|
110
|
+
for natrc in (Path.home() / ".natrc", Path(".natrc")):
|
|
111
|
+
_remove_path(natrc, removed, errors)
|
|
112
|
+
else:
|
|
113
|
+
print(" ℹ️ Skipping .natrc removal (--keep-config).")
|
|
114
|
+
|
|
115
|
+
# -----------------------------------------------------------------------
|
|
116
|
+
# Remove ~/.nat/ data directory
|
|
117
|
+
# -----------------------------------------------------------------------
|
|
118
|
+
nat_data_dir = Path.home() / ".nat"
|
|
119
|
+
if nat_data_dir.exists():
|
|
120
|
+
_remove_path(nat_data_dir, removed, errors)
|
|
121
|
+
|
|
122
|
+
# -----------------------------------------------------------------------
|
|
123
|
+
# Remove shell completion files
|
|
124
|
+
# -----------------------------------------------------------------------
|
|
125
|
+
for path in (*_BASH_COMPLETION_PATHS, *_ZSH_COMPLETION_PATHS, *_FISH_COMPLETION_PATHS):
|
|
126
|
+
if path.exists():
|
|
127
|
+
_remove_path(path, removed, errors)
|
|
128
|
+
|
|
129
|
+
# -----------------------------------------------------------------------
|
|
130
|
+
# Purge: alembic downgrade base
|
|
131
|
+
# -----------------------------------------------------------------------
|
|
132
|
+
if purge:
|
|
133
|
+
db_url = os.environ.get("DATABASE_URL")
|
|
134
|
+
if db_url:
|
|
135
|
+
print("\n 🗄️ Running alembic downgrade base …")
|
|
136
|
+
result = subprocess.run( # noqa: S603
|
|
137
|
+
["alembic", "downgrade", "base"],
|
|
138
|
+
check=False,
|
|
139
|
+
)
|
|
140
|
+
if result.returncode != 0:
|
|
141
|
+
print(" ❌ alembic downgrade failed.", file=sys.stderr)
|
|
142
|
+
errors.append("alembic downgrade base")
|
|
143
|
+
else:
|
|
144
|
+
print(" ✅ Database tables removed.")
|
|
145
|
+
else:
|
|
146
|
+
print(" ℹ️ DATABASE_URL not set — skipping alembic downgrade.")
|
|
147
|
+
|
|
148
|
+
# -----------------------------------------------------------------------
|
|
149
|
+
# Summary
|
|
150
|
+
# -----------------------------------------------------------------------
|
|
151
|
+
print()
|
|
152
|
+
if removed:
|
|
153
|
+
print(" ✅ Removed:")
|
|
154
|
+
for item in removed:
|
|
155
|
+
print(f" {item}")
|
|
156
|
+
else:
|
|
157
|
+
print(" ℹ️ Nothing to remove.")
|
|
158
|
+
|
|
159
|
+
if errors:
|
|
160
|
+
print("\n ⚠️ Errors encountered:")
|
|
161
|
+
for err in errors:
|
|
162
|
+
print(f" {err}", file=sys.stderr)
|
|
163
|
+
|
|
164
|
+
print()
|
|
165
|
+
print(" 📦 To complete uninstallation, run:")
|
|
166
|
+
print(" pip uninstall nat-engine")
|
|
167
|
+
print()
|
|
168
|
+
|
|
169
|
+
return _EXIT_OK if not errors else _EXIT_ERROR
|
mannf/product/upgrade.py
ADDED
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
# Copyright (C) 2026 Brad Guider
|
|
2
|
+
# This file is part of NAT (Neural Agent Testing Framework).
|
|
3
|
+
# Licensed under the AGPL-3.0. See LICENSE for details.
|
|
4
|
+
# Commercial licensing available — see COMMERCIAL_LICENSE.md.
|
|
5
|
+
|
|
6
|
+
"""Self-upgrade command for NAT — checks PyPI and upgrades the installed package.
|
|
7
|
+
|
|
8
|
+
Entry point::
|
|
9
|
+
|
|
10
|
+
def run_upgrade(args: argparse.Namespace) -> int:
|
|
11
|
+
...
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import argparse
|
|
17
|
+
import subprocess
|
|
18
|
+
import sys
|
|
19
|
+
|
|
20
|
+
_PYPI_JSON_URL = "https://pypi.org/pypi/nat-engine/json"
|
|
21
|
+
|
|
22
|
+
# ---------------------------------------------------------------------------
|
|
23
|
+
# Exit codes (reuse CLI constants pattern)
|
|
24
|
+
# ---------------------------------------------------------------------------
|
|
25
|
+
_EXIT_OK = 0
|
|
26
|
+
_EXIT_ERROR = 1
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _get_latest_pypi_version(allow_pre: bool = False) -> str | None:
|
|
30
|
+
"""Query PyPI JSON API and return the latest version string.
|
|
31
|
+
|
|
32
|
+
Parameters
|
|
33
|
+
----------
|
|
34
|
+
allow_pre:
|
|
35
|
+
When ``True``, consider pre-release versions as well.
|
|
36
|
+
|
|
37
|
+
Returns
|
|
38
|
+
-------
|
|
39
|
+
str | None
|
|
40
|
+
Latest version string, or ``None`` if the query fails.
|
|
41
|
+
"""
|
|
42
|
+
try:
|
|
43
|
+
import httpx # noqa: PLC0415
|
|
44
|
+
|
|
45
|
+
response = httpx.get(_PYPI_JSON_URL, timeout=10)
|
|
46
|
+
response.raise_for_status()
|
|
47
|
+
data = response.json()
|
|
48
|
+
if allow_pre:
|
|
49
|
+
# Return the overall latest release (may be pre-release)
|
|
50
|
+
versions = list(data.get("releases", {}).keys())
|
|
51
|
+
if not versions:
|
|
52
|
+
return None
|
|
53
|
+
# Sort by packaging.version if available, otherwise lexicographic
|
|
54
|
+
try:
|
|
55
|
+
from packaging.version import Version # noqa: PLC0415
|
|
56
|
+
|
|
57
|
+
versions.sort(key=Version)
|
|
58
|
+
except Exception:
|
|
59
|
+
versions.sort()
|
|
60
|
+
return versions[-1]
|
|
61
|
+
else:
|
|
62
|
+
return data.get("info", {}).get("version")
|
|
63
|
+
except Exception as exc: # noqa: BLE001
|
|
64
|
+
print(f" ⚠️ Could not query PyPI: {exc}", file=sys.stderr)
|
|
65
|
+
return None
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def run_upgrade(args: argparse.Namespace) -> int:
|
|
69
|
+
"""Check for updates and optionally upgrade ``nat-engine``.
|
|
70
|
+
|
|
71
|
+
Parameters
|
|
72
|
+
----------
|
|
73
|
+
args:
|
|
74
|
+
Parsed CLI arguments. The following attributes are read:
|
|
75
|
+
|
|
76
|
+
- ``check`` (bool) — only check for an update, do not install.
|
|
77
|
+
- ``migrate`` (bool) — run ``alembic upgrade head`` after upgrading.
|
|
78
|
+
- ``pre`` (bool) — allow pre-release versions when checking / installing.
|
|
79
|
+
|
|
80
|
+
Returns
|
|
81
|
+
-------
|
|
82
|
+
int
|
|
83
|
+
Exit code (0 = success, non-zero = error).
|
|
84
|
+
"""
|
|
85
|
+
from mannf._version import __version__ as current_version # noqa: PLC0415
|
|
86
|
+
|
|
87
|
+
check_only: bool = getattr(args, "check", False)
|
|
88
|
+
run_migrate: bool = getattr(args, "migrate", False)
|
|
89
|
+
allow_pre: bool = getattr(args, "pre", False)
|
|
90
|
+
|
|
91
|
+
print()
|
|
92
|
+
print(f" 🔍 Current version : nat-engine {current_version}")
|
|
93
|
+
|
|
94
|
+
latest = _get_latest_pypi_version(allow_pre=allow_pre)
|
|
95
|
+
if latest is None:
|
|
96
|
+
print(" ❌ Unable to determine the latest version from PyPI.", file=sys.stderr)
|
|
97
|
+
return _EXIT_ERROR
|
|
98
|
+
|
|
99
|
+
print(f" 🌐 Latest version : nat-engine {latest}")
|
|
100
|
+
|
|
101
|
+
if current_version == latest:
|
|
102
|
+
print(" ✅ You are already on the latest version.")
|
|
103
|
+
print()
|
|
104
|
+
return _EXIT_OK
|
|
105
|
+
|
|
106
|
+
if check_only:
|
|
107
|
+
print(f" ℹ️ An update is available: {current_version} → {latest}")
|
|
108
|
+
print(" Run `nat upgrade` (without --check) to install it.")
|
|
109
|
+
print()
|
|
110
|
+
return _EXIT_OK
|
|
111
|
+
|
|
112
|
+
# -----------------------------------------------------------------------
|
|
113
|
+
# Perform the upgrade
|
|
114
|
+
# -----------------------------------------------------------------------
|
|
115
|
+
pip_cmd = [sys.executable, "-m", "pip", "install", "--upgrade", "nat-engine"]
|
|
116
|
+
if allow_pre:
|
|
117
|
+
pip_cmd.append("--pre")
|
|
118
|
+
|
|
119
|
+
print(f"\n ⬆️ Upgrading nat-engine {current_version} → {latest} …")
|
|
120
|
+
result = subprocess.run(pip_cmd, check=False) # noqa: S603
|
|
121
|
+
if result.returncode != 0:
|
|
122
|
+
print(" ❌ pip upgrade failed.", file=sys.stderr)
|
|
123
|
+
return _EXIT_ERROR
|
|
124
|
+
|
|
125
|
+
# -----------------------------------------------------------------------
|
|
126
|
+
# Optional DB migration
|
|
127
|
+
# -----------------------------------------------------------------------
|
|
128
|
+
if run_migrate:
|
|
129
|
+
print("\n 🗄️ Running database migrations (alembic upgrade head) …")
|
|
130
|
+
alembic_cmd = ["alembic", "upgrade", "head"]
|
|
131
|
+
migrate_result = subprocess.run(alembic_cmd, check=False) # noqa: S603
|
|
132
|
+
if migrate_result.returncode != 0:
|
|
133
|
+
print(" ❌ alembic migration failed.", file=sys.stderr)
|
|
134
|
+
return _EXIT_ERROR
|
|
135
|
+
print(" ✅ Database migrations complete.")
|
|
136
|
+
|
|
137
|
+
print(f"\n 🎉 nat-engine upgraded: {current_version} → {latest}")
|
|
138
|
+
print()
|
|
139
|
+
return _EXIT_OK
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
# Copyright (C) 2026 Brad Guider
|
|
2
|
+
# This file is part of NAT (Neural Agent Testing Framework).
|
|
3
|
+
# Licensed under the AGPL-3.0. See LICENSE for details.
|
|
4
|
+
# Commercial licensing available — see COMMERCIAL_LICENSE.md.
|
|
5
|
+
|
|
6
|
+
"""Weight persistence and transfer learning package."""
|
|
7
|
+
|
|
8
|
+
from mannf.product.weights.store import WeightStore
|
|
9
|
+
from mannf.product.weights.registry import WeightRegistry
|
|
10
|
+
from mannf.product.weights.blob_store import BlobWeightStore
|
|
11
|
+
from mannf.product.weights.factory import get_weight_store
|
|
12
|
+
|
|
13
|
+
__all__ = ["WeightStore", "WeightRegistry", "BlobWeightStore", "get_weight_store"]
|
|
@@ -0,0 +1,299 @@
|
|
|
1
|
+
# Copyright (C) 2026 Brad Guider
|
|
2
|
+
# This file is part of NAT (Neural Agent Testing Framework).
|
|
3
|
+
# Licensed under the AGPL-3.0. See LICENSE for details.
|
|
4
|
+
# Commercial licensing available — see COMMERCIAL_LICENSE.md.
|
|
5
|
+
|
|
6
|
+
"""Azure Blob Storage backend for weight persistence.
|
|
7
|
+
|
|
8
|
+
Activated when ``AZURE_STORAGE_CONNECTION_STRING`` or
|
|
9
|
+
``AZURE_STORAGE_ACCOUNT_NAME`` is set in the environment. Falls back to
|
|
10
|
+
the local :class:`~mannf.product.weights.store.WeightStore` when neither
|
|
11
|
+
variable is present.
|
|
12
|
+
|
|
13
|
+
The blob JSON format is identical to the on-disk format used by
|
|
14
|
+
:class:`~mannf.product.weights.store.WeightStore` so the two backends are
|
|
15
|
+
fully interchangeable.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import json
|
|
21
|
+
import logging
|
|
22
|
+
import os
|
|
23
|
+
from datetime import datetime, timezone
|
|
24
|
+
from typing import Any, Dict, List
|
|
25
|
+
|
|
26
|
+
import numpy as np
|
|
27
|
+
|
|
28
|
+
logger = logging.getLogger(__name__)
|
|
29
|
+
|
|
30
|
+
_FORMAT_VERSION = "1.0"
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _layers_to_serialisable(layers: List[dict]) -> List[dict]:
|
|
34
|
+
"""Convert ``{"W": ndarray, "b": ndarray}`` layer dicts to JSON-safe lists."""
|
|
35
|
+
return [{"W": layer["W"].tolist(), "b": layer["b"].tolist()} for layer in layers]
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _layers_from_serialisable(layers: List[dict]) -> List[dict]:
|
|
39
|
+
"""Restore layer dicts from JSON-loaded lists back to ``np.ndarray`` values."""
|
|
40
|
+
return [
|
|
41
|
+
{"W": np.array(layer["W"], dtype=float), "b": np.array(layer["b"], dtype=float)}
|
|
42
|
+
for layer in layers
|
|
43
|
+
]
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class BlobWeightStore:
|
|
47
|
+
"""Azure Blob Storage backend for weight snapshots.
|
|
48
|
+
|
|
49
|
+
Activated when ``AZURE_STORAGE_CONNECTION_STRING`` or
|
|
50
|
+
``AZURE_STORAGE_ACCOUNT_NAME`` is set. Falls back to the local
|
|
51
|
+
:class:`~mannf.product.weights.store.WeightStore` when not configured.
|
|
52
|
+
|
|
53
|
+
Parameters
|
|
54
|
+
----------
|
|
55
|
+
container_name:
|
|
56
|
+
Name of the Azure Blob Storage container to use. Defaults to
|
|
57
|
+
``"nat-weights"``. The container is created automatically on first
|
|
58
|
+
use if it does not already exist.
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
def __init__(self, container_name: str = "nat-weights") -> None:
|
|
62
|
+
self._container_name = container_name
|
|
63
|
+
self._client = self._build_client()
|
|
64
|
+
|
|
65
|
+
# ------------------------------------------------------------------
|
|
66
|
+
# Public async API
|
|
67
|
+
# ------------------------------------------------------------------
|
|
68
|
+
|
|
69
|
+
async def save(self, blob_path: str, agent_weights: Dict[str, Any]) -> str:
|
|
70
|
+
"""Upload weight JSON to blob storage.
|
|
71
|
+
|
|
72
|
+
Parameters
|
|
73
|
+
----------
|
|
74
|
+
blob_path:
|
|
75
|
+
The blob name / path inside the container (e.g.
|
|
76
|
+
``"tenant-a/my-snapshot.json"``).
|
|
77
|
+
agent_weights:
|
|
78
|
+
Mapping of ``agent_id`` → ``{"role": str, "layers": List[dict]}``.
|
|
79
|
+
Each layer dict must have ``"W"`` and ``"b"`` numpy arrays.
|
|
80
|
+
|
|
81
|
+
Returns
|
|
82
|
+
-------
|
|
83
|
+
str
|
|
84
|
+
The URL of the uploaded blob.
|
|
85
|
+
"""
|
|
86
|
+
agents_serialisable: Dict[str, Any] = {}
|
|
87
|
+
for agent_id, info in agent_weights.items():
|
|
88
|
+
agents_serialisable[agent_id] = {
|
|
89
|
+
"role": info.get("role", ""),
|
|
90
|
+
"layers": _layers_to_serialisable(info["layers"]),
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
payload = {
|
|
94
|
+
"version": _FORMAT_VERSION,
|
|
95
|
+
"saved_at": datetime.now(timezone.utc).isoformat(),
|
|
96
|
+
"agents": agents_serialisable,
|
|
97
|
+
}
|
|
98
|
+
data = json.dumps(payload, indent=2).encode("utf-8")
|
|
99
|
+
|
|
100
|
+
blob_client = self._client.get_blob_client(
|
|
101
|
+
container=self._container_name, blob=blob_path
|
|
102
|
+
)
|
|
103
|
+
await self._ensure_container()
|
|
104
|
+
await blob_client.upload_blob(data, overwrite=True)
|
|
105
|
+
url = blob_client.url
|
|
106
|
+
logger.info("BlobWeightStore: uploaded %s (%d agents)", blob_path, len(agent_weights))
|
|
107
|
+
return url
|
|
108
|
+
|
|
109
|
+
async def load(self, blob_path: str) -> Dict[str, Any]:
|
|
110
|
+
"""Download and parse weight JSON from blob storage.
|
|
111
|
+
|
|
112
|
+
Parameters
|
|
113
|
+
----------
|
|
114
|
+
blob_path:
|
|
115
|
+
The blob name / path inside the container.
|
|
116
|
+
|
|
117
|
+
Returns
|
|
118
|
+
-------
|
|
119
|
+
dict
|
|
120
|
+
``{"version": str, "saved_at": str, "agents": {agent_id: {...}}}``
|
|
121
|
+
where each layer dict contains ``np.ndarray`` values.
|
|
122
|
+
|
|
123
|
+
Raises
|
|
124
|
+
------
|
|
125
|
+
FileNotFoundError
|
|
126
|
+
If the blob does not exist.
|
|
127
|
+
ValueError
|
|
128
|
+
If the blob content cannot be parsed or fails validation.
|
|
129
|
+
"""
|
|
130
|
+
from azure.core.exceptions import ResourceNotFoundError
|
|
131
|
+
|
|
132
|
+
blob_client = self._client.get_blob_client(
|
|
133
|
+
container=self._container_name, blob=blob_path
|
|
134
|
+
)
|
|
135
|
+
try:
|
|
136
|
+
stream = await blob_client.download_blob()
|
|
137
|
+
raw = await stream.readall()
|
|
138
|
+
except ResourceNotFoundError as exc:
|
|
139
|
+
raise FileNotFoundError(f"Weight blob not found: {blob_path!r}") from exc
|
|
140
|
+
|
|
141
|
+
try:
|
|
142
|
+
data = json.loads(raw)
|
|
143
|
+
except json.JSONDecodeError as exc:
|
|
144
|
+
raise ValueError(f"Corrupt weight blob {blob_path!r}: {exc}") from exc
|
|
145
|
+
|
|
146
|
+
_validate_structure(data)
|
|
147
|
+
|
|
148
|
+
agents_out: Dict[str, Any] = {}
|
|
149
|
+
for agent_id, info in data["agents"].items():
|
|
150
|
+
agents_out[agent_id] = {
|
|
151
|
+
"role": info.get("role", ""),
|
|
152
|
+
"layers": _layers_from_serialisable(info["layers"]),
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
return {
|
|
156
|
+
"version": data["version"],
|
|
157
|
+
"saved_at": data["saved_at"],
|
|
158
|
+
"agents": agents_out,
|
|
159
|
+
}
|
|
160
|
+
|
|
161
|
+
async def delete(self, blob_path: str) -> None:
|
|
162
|
+
"""Delete a weight blob.
|
|
163
|
+
|
|
164
|
+
Parameters
|
|
165
|
+
----------
|
|
166
|
+
blob_path:
|
|
167
|
+
The blob name / path inside the container.
|
|
168
|
+
|
|
169
|
+
Raises
|
|
170
|
+
------
|
|
171
|
+
FileNotFoundError
|
|
172
|
+
If the blob does not exist.
|
|
173
|
+
"""
|
|
174
|
+
from azure.core.exceptions import ResourceNotFoundError
|
|
175
|
+
|
|
176
|
+
blob_client = self._client.get_blob_client(
|
|
177
|
+
container=self._container_name, blob=blob_path
|
|
178
|
+
)
|
|
179
|
+
try:
|
|
180
|
+
await blob_client.delete_blob()
|
|
181
|
+
except ResourceNotFoundError as exc:
|
|
182
|
+
raise FileNotFoundError(f"Weight blob not found: {blob_path!r}") from exc
|
|
183
|
+
logger.info("BlobWeightStore: deleted %s", blob_path)
|
|
184
|
+
|
|
185
|
+
async def list_blobs(self, prefix: str = "") -> List[str]:
|
|
186
|
+
"""List weight blobs with an optional prefix filter.
|
|
187
|
+
|
|
188
|
+
Parameters
|
|
189
|
+
----------
|
|
190
|
+
prefix:
|
|
191
|
+
If non-empty, only blobs whose names begin with *prefix* are
|
|
192
|
+
returned.
|
|
193
|
+
|
|
194
|
+
Returns
|
|
195
|
+
-------
|
|
196
|
+
list[str]
|
|
197
|
+
Blob names (paths) within the container.
|
|
198
|
+
"""
|
|
199
|
+
container_client = self._client.get_container_client(self._container_name)
|
|
200
|
+
blobs: List[str] = []
|
|
201
|
+
try:
|
|
202
|
+
async for blob in container_client.list_blobs(name_starts_with=prefix or None):
|
|
203
|
+
blobs.append(blob.name)
|
|
204
|
+
except Exception: # container may not exist yet
|
|
205
|
+
pass
|
|
206
|
+
return blobs
|
|
207
|
+
|
|
208
|
+
# ------------------------------------------------------------------
|
|
209
|
+
# Class-level helpers
|
|
210
|
+
# ------------------------------------------------------------------
|
|
211
|
+
|
|
212
|
+
@staticmethod
|
|
213
|
+
def is_configured() -> bool:
|
|
214
|
+
"""Return ``True`` when Azure Blob Storage environment variables are set."""
|
|
215
|
+
return bool(
|
|
216
|
+
os.environ.get("AZURE_STORAGE_CONNECTION_STRING")
|
|
217
|
+
or os.environ.get("AZURE_STORAGE_ACCOUNT_NAME")
|
|
218
|
+
)
|
|
219
|
+
|
|
220
|
+
# ------------------------------------------------------------------
|
|
221
|
+
# Internal helpers
|
|
222
|
+
# ------------------------------------------------------------------
|
|
223
|
+
|
|
224
|
+
def _build_client(self): # type: ignore[return]
|
|
225
|
+
"""Construct an ``AsyncBlobServiceClient`` from environment variables."""
|
|
226
|
+
try:
|
|
227
|
+
from azure.storage.blob.aio import BlobServiceClient
|
|
228
|
+
except ImportError as exc:
|
|
229
|
+
raise ImportError(
|
|
230
|
+
"azure-storage-blob is required for BlobWeightStore. "
|
|
231
|
+
"Install it with: pip install 'mannf[azure-storage]'"
|
|
232
|
+
) from exc
|
|
233
|
+
|
|
234
|
+
conn_str = os.environ.get("AZURE_STORAGE_CONNECTION_STRING")
|
|
235
|
+
if conn_str:
|
|
236
|
+
return BlobServiceClient.from_connection_string(conn_str)
|
|
237
|
+
|
|
238
|
+
account_name = os.environ.get("AZURE_STORAGE_ACCOUNT_NAME")
|
|
239
|
+
if account_name:
|
|
240
|
+
account_key = os.environ.get("AZURE_STORAGE_ACCOUNT_KEY", "")
|
|
241
|
+
if account_key:
|
|
242
|
+
conn_str = (
|
|
243
|
+
f"DefaultEndpointsProtocol=https;"
|
|
244
|
+
f"AccountName={account_name};"
|
|
245
|
+
f"AccountKey={account_key};"
|
|
246
|
+
f"EndpointSuffix=core.windows.net"
|
|
247
|
+
)
|
|
248
|
+
return BlobServiceClient.from_connection_string(conn_str)
|
|
249
|
+
# Managed identity / passwordless
|
|
250
|
+
from azure.identity.aio import DefaultAzureCredential
|
|
251
|
+
credential = DefaultAzureCredential()
|
|
252
|
+
return BlobServiceClient(
|
|
253
|
+
account_url=f"https://{account_name}.blob.core.windows.net",
|
|
254
|
+
credential=credential,
|
|
255
|
+
)
|
|
256
|
+
|
|
257
|
+
raise EnvironmentError(
|
|
258
|
+
"Neither AZURE_STORAGE_CONNECTION_STRING nor "
|
|
259
|
+
"AZURE_STORAGE_ACCOUNT_NAME is set."
|
|
260
|
+
)
|
|
261
|
+
|
|
262
|
+
async def _ensure_container(self) -> None:
|
|
263
|
+
"""Create the container if it does not already exist."""
|
|
264
|
+
from azure.core.exceptions import ResourceExistsError
|
|
265
|
+
|
|
266
|
+
container_client = self._client.get_container_client(self._container_name)
|
|
267
|
+
try:
|
|
268
|
+
await container_client.create_container()
|
|
269
|
+
logger.debug("BlobWeightStore: created container %r", self._container_name)
|
|
270
|
+
except ResourceExistsError:
|
|
271
|
+
pass
|
|
272
|
+
|
|
273
|
+
|
|
274
|
+
# ---------------------------------------------------------------------------
|
|
275
|
+
# Shared validation helper (mirrors WeightStore._validate_structure)
|
|
276
|
+
# ---------------------------------------------------------------------------
|
|
277
|
+
|
|
278
|
+
|
|
279
|
+
def _validate_structure(data: Any) -> None:
|
|
280
|
+
"""Raise :class:`ValueError` if *data* does not match the expected schema."""
|
|
281
|
+
if not isinstance(data, dict):
|
|
282
|
+
raise ValueError("Weight blob root must be a JSON object")
|
|
283
|
+
if data.get("version") != _FORMAT_VERSION:
|
|
284
|
+
raise ValueError(
|
|
285
|
+
f"Unsupported weight blob version {data.get('version')!r} "
|
|
286
|
+
f"(expected {_FORMAT_VERSION!r})"
|
|
287
|
+
)
|
|
288
|
+
if "agents" not in data or not isinstance(data["agents"], dict):
|
|
289
|
+
raise ValueError("Weight blob missing 'agents' mapping")
|
|
290
|
+
for agent_id, info in data["agents"].items():
|
|
291
|
+
if not isinstance(info, dict):
|
|
292
|
+
raise ValueError(f"Agent entry {agent_id!r} is not a dict")
|
|
293
|
+
if "layers" not in info or not isinstance(info["layers"], list):
|
|
294
|
+
raise ValueError(f"Agent {agent_id!r} missing 'layers' list")
|
|
295
|
+
for i, layer in enumerate(info["layers"]):
|
|
296
|
+
if not isinstance(layer, dict) or "W" not in layer or "b" not in layer:
|
|
297
|
+
raise ValueError(
|
|
298
|
+
f"Agent {agent_id!r} layer {i} must have 'W' and 'b' keys"
|
|
299
|
+
)
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
# Copyright (C) 2026 Brad Guider
|
|
2
|
+
# This file is part of NAT (Neural Agent Testing Framework).
|
|
3
|
+
# Licensed under the AGPL-3.0. See LICENSE for details.
|
|
4
|
+
# Commercial licensing available — see COMMERCIAL_LICENSE.md.
|
|
5
|
+
|
|
6
|
+
"""Factory for selecting the appropriate weight-store backend.
|
|
7
|
+
|
|
8
|
+
Call :func:`get_weight_store` to obtain a store instance. The function
|
|
9
|
+
returns a :class:`~mannf.product.weights.blob_store.BlobWeightStore` when
|
|
10
|
+
Azure Blob Storage environment variables are present, or the local
|
|
11
|
+
:class:`~mannf.product.weights.store.WeightStore` otherwise.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
from typing import Union
|
|
17
|
+
|
|
18
|
+
from mannf.product.weights.blob_store import BlobWeightStore
|
|
19
|
+
from mannf.product.weights.store import WeightStore
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def get_weight_store(container_name: str = "nat-weights") -> Union[BlobWeightStore, WeightStore]:
|
|
23
|
+
"""Return the appropriate weight-store backend.
|
|
24
|
+
|
|
25
|
+
If ``AZURE_STORAGE_CONNECTION_STRING`` or ``AZURE_STORAGE_ACCOUNT_NAME``
|
|
26
|
+
is set in the environment, a :class:`BlobWeightStore` is returned.
|
|
27
|
+
Otherwise the local :class:`WeightStore` is returned (default behaviour,
|
|
28
|
+
fully preserves existing local-file semantics).
|
|
29
|
+
|
|
30
|
+
Parameters
|
|
31
|
+
----------
|
|
32
|
+
container_name:
|
|
33
|
+
Passed through to :class:`BlobWeightStore` when Azure is configured.
|
|
34
|
+
Ignored for the local backend.
|
|
35
|
+
|
|
36
|
+
Returns
|
|
37
|
+
-------
|
|
38
|
+
BlobWeightStore | WeightStore
|
|
39
|
+
"""
|
|
40
|
+
if BlobWeightStore.is_configured():
|
|
41
|
+
return BlobWeightStore(container_name=container_name)
|
|
42
|
+
return WeightStore()
|