[haik]: refactor: systematize naming conventions across 20+ backend modules. Private module-level vars renamed _foo to P_FOO, private funcs _foo to p_foo, cross-module public funcs drop leading underscores entirely (e.g. _discover_mcp_tools_http -> discover_mcp_tools_http, _load_all -> load_all_tools, _classify_services -> classify_services). Local vars inside functions lose underscores (_backend -> backend_dir, _port -> port). Unnecessary import aliases removed (import json as _json -> import json). Imports reorganized to proper source modules (derive_mcp_config from mcp_config, oauth refresh funcs from oauth_tokens). Added Public comments documenting cross-module API boundaries. Extracted PENDING_OAUTH to oauth_state.py to break circular import. No behavioral changes.

This commit is contained in:
haikdc
2026-06-13 23:38:39 -07:00
parent 67fce8ee89
commit 08eb4b561f
23 changed files with 530 additions and 501 deletions
+7 -6
View File
@@ -14,17 +14,18 @@ from backend.apps.agents.core.models import (
)
from backend.apps.agents.core.ws_manager import ws_manager
from backend.apps.settings.store import load_settings
from backend.apps.tools_lib.oauth_tokens import (
refresh_google_token,
refresh_airtable_token,
refresh_hubspot_token,
)
from backend.apps.tools_lib.tools_lib import (
_load_all as load_all_tools,
derive_mcp_config,
load_all_tools,
load_builtin_permissions,
load_trusted_sensitive_paths,
refresh_airtable_token,
refresh_google_token,
refresh_hubspot_token,
save_trusted_sensitive_paths,
)
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name, derive_mcp_config
from backend.apps.agents.core.error_classify import (
is_auth_error,
is_free_trial_exhausted,
+2 -2
View File
@@ -373,8 +373,8 @@ async def subscriptions_connect(body: dict):
result = await start_oauth(provider)
if result.get("flow") == "authorization_code" and result.get("state"):
from backend.main import _pending_oauth
_pending_oauth[result["state"]] = {
from backend.apps.oauth_state import PENDING_OAUTH
PENDING_OAUTH[result["state"]] = {
"provider": provider,
"code_verifier": result.get("code_verifier", ""),
"redirect_uri": result.get("redirect_uri", ""),
+1 -1
View File
@@ -11,7 +11,7 @@ from typing import Any
from backend.apps.agents.providers.registry import resolve_aux_model
from backend.apps.settings.credentials import get_anthropic_client_for_model
from backend.apps.settings.store import load_settings
from backend.apps.tools_lib.tools_lib import _load_all as load_all_tools
from backend.apps.tools_lib.tools_lib import load_all_tools
logger = logging.getLogger(__name__)
@@ -2,7 +2,7 @@ from typing import Callable
from backend.apps.modes.modes import load_mode
from backend.apps.tools_lib.tools_lib import (
_load_all as load_all_tools,
load_all_tools,
)
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
from backend.apps.agents.manager.prompt.tool_catalog import get_denied_tool_names, is_fully_denied
+29 -29
View File
@@ -19,10 +19,10 @@ from backend.apps.settings.models import AppSettings, DEFAULT_SYSTEM_PROMPT
logger = logging.getLogger(__name__)
SETTINGS_FILE = os.path.join(DATA_DIR, "settings.json")
P_SETTINGS_FILE = os.path.join(DATA_DIR, "settings.json")
def _migrate_legacy_fields(raw: dict) -> dict:
def p_migrate_legacy_fields(raw: dict) -> dict:
"""Translate deprecated pre-launch field names ('managed', 'openswarm_auth_token') into production schema."""
if raw.get("connection_mode") == "managed":
raw["connection_mode"] = "openswarm-pro"
@@ -31,7 +31,7 @@ def _migrate_legacy_fields(raw: dict) -> dict:
return raw
def _coerce_settings(raw: dict) -> AppSettings:
def p_coerce_settings(raw: dict) -> AppSettings:
"""Build AppSettings, surviving a settings.json written by a different app
version. Unknown fields are already ignored by pydantic; the case this guards
is a field whose TYPE drifted across versions (e.g. a list that is now a
@@ -55,12 +55,12 @@ def _coerce_settings(raw: dict) -> AppSettings:
return AppSettings()
def _preserve_corrupt_settings() -> None:
def p_preserve_corrupt_settings() -> None:
"""Move an unparseable settings.json aside so boot proceeds on defaults while
the original stays recoverable (the next save would otherwise overwrite it)."""
try:
backup = SETTINGS_FILE + ".corrupt"
os.replace(SETTINGS_FILE, backup)
backup = P_SETTINGS_FILE + ".corrupt"
os.replace(P_SETTINGS_FILE, backup)
logger.warning("settings.json was unparseable; preserved at %s", backup)
except OSError:
pass
@@ -70,13 +70,13 @@ def _preserve_corrupt_settings() -> None:
# so even a hand-edited file or an unexpected writer is picked up immediately. A stat
# skips the open+parse+validate that Defender turns into 5-50ms on Windows. Copies on
# both sides keep handler isolation: callers mutate their copy, never the cache.
_cached_settings: AppSettings | None = None
_cached_sig: tuple[int, int] | None = None
P_CACHED_SETTINGS: AppSettings | None = None
P_CACHED_SIG: tuple[int, int] | None = None
def _settings_sig() -> tuple[int, int] | None:
def p_settings_sig() -> tuple[int, int] | None:
try:
st = os.stat(SETTINGS_FILE)
st = os.stat(P_SETTINGS_FILE)
return (st.st_mtime_ns, st.st_size)
except OSError:
return None
@@ -86,38 +86,38 @@ def load_settings() -> AppSettings:
"""Load settings from JSON file, returning defaults if not found. Never raises
on a corrupt or version-mismatched file: a single bad settings.json must not
brick boot (it is read at startup, by the settings endpoint, and per dispatch)."""
global _cached_settings, _cached_sig
sig = _settings_sig()
if sig is not None and _cached_settings is not None and sig == _cached_sig:
return _cached_settings.model_copy(deep=True)
if os.path.exists(SETTINGS_FILE):
global P_CACHED_SETTINGS, P_CACHED_SIG
sig = p_settings_sig()
if sig is not None and P_CACHED_SETTINGS is not None and sig == P_CACHED_SIG:
return P_CACHED_SETTINGS.model_copy(deep=True)
if os.path.exists(P_SETTINGS_FILE):
try:
with open(SETTINGS_FILE) as f:
with open(P_SETTINGS_FILE) as f:
raw = json.load(f)
except (json.JSONDecodeError, OSError, ValueError):
_preserve_corrupt_settings()
p_preserve_corrupt_settings()
return AppSettings()
if not isinstance(raw, dict):
# Valid JSON but not an object (e.g. a bare list/number); unusable.
_preserve_corrupt_settings()
p_preserve_corrupt_settings()
return AppSettings()
settings = _coerce_settings(_migrate_legacy_fields(raw))
settings = p_coerce_settings(p_migrate_legacy_fields(raw))
if settings.default_system_prompt is None:
settings.default_system_prompt = DEFAULT_SYSTEM_PROMPT
_cached_settings = settings.model_copy(deep=True)
_cached_sig = sig
P_CACHED_SETTINGS = settings.model_copy(deep=True)
P_CACHED_SIG = sig
return settings
return AppSettings()
# threading.Lock guards every SETTINGS_FILE write; works for sync paths and async run_in_executor paths.
_settings_write_lock = threading.Lock()
P_SETTINGS_WRITE_LOCK = threading.Lock()
def atomic_write_settings(payload: dict) -> None:
"""Atomic SETTINGS_FILE write; call via save_settings*, not directly."""
global _cached_settings, _cached_sig
with _settings_write_lock:
global P_CACHED_SETTINGS, P_CACHED_SIG
with P_SETTINGS_WRITE_LOCK:
os.makedirs(DATA_DIR, exist_ok=True)
fd, tmp = tempfile.mkstemp(prefix=".settings.", suffix=".tmp", dir=DATA_DIR)
try:
@@ -126,12 +126,12 @@ def atomic_write_settings(payload: dict) -> None:
# Windows: Defender can briefly lock the destination; one retry handles every real case.
for attempt in range(2):
try:
os.replace(tmp, SETTINGS_FILE)
os.replace(tmp, P_SETTINGS_FILE)
# Refresh the cache inside the lock so cache order matches disk order.
_cached_settings = _coerce_settings(_migrate_legacy_fields(dict(payload)))
if _cached_settings.default_system_prompt is None:
_cached_settings.default_system_prompt = DEFAULT_SYSTEM_PROMPT
_cached_sig = _settings_sig()
P_CACHED_SETTINGS = p_coerce_settings(p_migrate_legacy_fields(dict(payload)))
if P_CACHED_SETTINGS.default_system_prompt is None:
P_CACHED_SETTINGS.default_system_prompt = DEFAULT_SYSTEM_PROMPT
P_CACHED_SIG = p_settings_sig()
return
except PermissionError:
if attempt == 1:
+35 -35
View File
@@ -11,19 +11,19 @@ from backend.config.Apps import SubApp
logger = logging.getLogger(__name__)
REPO = "anthropics/skills"
BRANCH = "main"
RAW_BASE = f"https://raw.githubusercontent.com/{REPO}/{BRANCH}"
MANIFEST_URL = f"{RAW_BASE}/.claude-plugin/marketplace.json"
REFRESH_INTERVAL_S = 3600
CONCURRENT_FETCHES = 15
P_REPO = "anthropics/skills"
P_BRANCH = "main"
P_RAW_BASE = f"https://raw.githubusercontent.com/{P_REPO}/{P_BRANCH}"
P_MANIFEST_URL = f"{P_RAW_BASE}/.claude-plugin/marketplace.json"
P_REFRESH_INTERVAL_S = 3600
P_CONCURRENT_FETCHES = 15
_cache: dict[str, dict] = {}
_cache_updated_at: float = 0
_refresh_task: Optional[asyncio.Task] = None
P_CACHE: dict[str, dict] = {}
P_CACHE_UPDATED_AT: float = 0
P_REFRESH_TASK: Optional[asyncio.Task] = None
def _parse_frontmatter(raw: str) -> tuple[dict, str]:
def p_parse_frontmatter(raw: str) -> tuple[dict, str]:
"""Split YAML frontmatter from markdown body."""
if not raw.startswith("---"):
return {}, raw
@@ -40,12 +40,12 @@ def _parse_frontmatter(raw: str) -> tuple[dict, str]:
return meta, body
async def _fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]:
async def p_fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]:
"""Fetch the marketplace.json manifest and return (skill_folder, plugin_name) pairs.
Uses raw.githubusercontent.com; no GitHub API needed, no rate limiting.
"""
resp = await client.get(MANIFEST_URL)
resp = await client.get(P_MANIFEST_URL)
resp.raise_for_status()
manifest = resp.json()
@@ -58,7 +58,7 @@ async def _fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]
return paths
async def _fetch_one_skill(
async def p_fetch_one_skill(
client: httpx.AsyncClient,
sem: asyncio.Semaphore,
folder: str,
@@ -66,7 +66,7 @@ async def _fetch_one_skill(
) -> Optional[dict]:
async with sem:
try:
resp = await client.get(f"{RAW_BASE}/{folder}/SKILL.md")
resp = await client.get(f"{P_RAW_BASE}/{folder}/SKILL.md")
if resp.status_code != 200:
return None
raw = resp.text
@@ -74,7 +74,7 @@ async def _fetch_one_skill(
logger.debug(f"Failed to fetch {folder}/SKILL.md: {exc}")
return None
meta, body = _parse_frontmatter(raw)
meta, body = p_parse_frontmatter(raw)
name = meta.get("name", "")
if not name:
folder_name = folder.rsplit("/", 1)[-1]
@@ -86,23 +86,23 @@ async def _fetch_one_skill(
"content": body,
"folder": folder,
"category": plugin_name.replace("-", " ").replace("_", " ").title(),
"repositoryUrl": f"https://github.com/{REPO}/tree/{BRANCH}/{folder}",
"repositoryUrl": f"https://github.com/{P_REPO}/tree/{P_BRANCH}/{folder}",
}
async def _fetch_all_skills() -> dict[str, dict]:
async def p_fetch_all_skills() -> dict[str, dict]:
skills: dict[str, dict] = {}
async with httpx.AsyncClient(timeout=30.0) as client:
try:
paths = await _fetch_skill_paths(client)
paths = await p_fetch_skill_paths(client)
except Exception as e:
logger.warning(f"Skill registry manifest fetch failed: {e}")
return skills
logger.info(f"Skill registry: found {len(paths)} skills in manifest, fetching content...")
sem = asyncio.Semaphore(CONCURRENT_FETCHES)
sem = asyncio.Semaphore(P_CONCURRENT_FETCHES)
results = await asyncio.gather(
*[_fetch_one_skill(client, sem, folder, plugin) for folder, plugin in paths]
*[p_fetch_one_skill(client, sem, folder, plugin) for folder, plugin in paths]
)
for rec in results:
if rec:
@@ -112,26 +112,26 @@ async def _fetch_all_skills() -> dict[str, dict]:
return skills
async def _refresh_loop():
global _cache, _cache_updated_at
async def p_refresh_loop():
global P_CACHE, P_CACHE_UPDATED_AT
while True:
try:
_cache = await _fetch_all_skills()
_cache_updated_at = time.time()
P_CACHE = await p_fetch_all_skills()
P_CACHE_UPDATED_AT = time.time()
except Exception as e:
logger.exception(f"Skill registry refresh error: {e}")
await asyncio.sleep(REFRESH_INTERVAL_S)
await asyncio.sleep(P_REFRESH_INTERVAL_S)
@asynccontextmanager
async def skill_registry_lifespan():
global _refresh_task
_refresh_task = asyncio.create_task(_refresh_loop())
global P_REFRESH_TASK
P_REFRESH_TASK = asyncio.create_task(p_refresh_loop())
yield
if _refresh_task:
_refresh_task.cancel()
if P_REFRESH_TASK:
P_REFRESH_TASK.cancel()
try:
await _refresh_task
await P_REFRESH_TASK
except asyncio.CancelledError:
pass
@@ -142,13 +142,13 @@ skill_registry = SubApp("skill-registry", skill_registry_lifespan)
@skill_registry.router.get("/stats")
async def registry_stats():
categories: dict[str, int] = {}
for s in _cache.values():
for s in P_CACHE.values():
cat = s.get("category", "General")
categories[cat] = categories.get(cat, 0) + 1
return {
"total": len(_cache),
"total": len(P_CACHE),
"categories": categories,
"lastUpdated": _cache_updated_at,
"lastUpdated": P_CACHE_UPDATED_AT,
}
@@ -159,7 +159,7 @@ async def registry_search(
offset: int = Query(0, ge=0),
category: str = Query("", description="Filter by category"),
):
pool = list(_cache.values())
pool = list(P_CACHE.values())
if category:
cat_lower = category.lower()
pool = [s for s in pool if s.get("category", "").lower() == cat_lower]
@@ -192,7 +192,7 @@ async def registry_search(
@skill_registry.router.get("/detail/{skill_name:path}")
async def registry_detail(skill_name: str):
sk = _cache.get(skill_name)
sk = P_CACHE.get(skill_name)
if not sk:
return {"error": "Skill not found"}, 404
return {"skill": sk}
+20 -20
View File
@@ -15,14 +15,14 @@ INDEX_PATH = os.path.join(SKILLS_DIR, ".skills_index.json")
from backend.config.paths import SKILLS_WORKSPACE_DIR
def _load_index() -> dict[str, dict]:
def p_load_index() -> dict[str, dict]:
if os.path.exists(INDEX_PATH):
with open(INDEX_PATH) as f:
return json.load(f)
return {}
def _save_index(index: dict[str, dict]):
def p_save_index(index: dict[str, dict]):
with open(INDEX_PATH, "w") as f:
json.dump(index, f, indent=2)
@@ -32,7 +32,7 @@ def _save_index(index: dict[str, dict]):
# `built_in: true` in the index. Users can edit the content (their
# changes flow through to the matching agent's prompt on the next turn),
# but they can't delete the file; the DELETE endpoint refuses with 409.
def _built_in_skill_registry() -> list[dict]:
def p_built_in_skill_registry() -> list[dict]:
# Imported lazily so this module stays cheap to import from
# everywhere (the skills outputs module pulls in pydantic+fastapi
# transitively and we don't want a cycle).
@@ -69,15 +69,15 @@ def _built_in_skill_registry() -> list[dict]:
]
def _seed_built_in_skills() -> None:
def p_seed_built_in_skills() -> None:
"""Copy each built-in skill into SKILLS_DIR if not already present, and
ensure the index has the `built_in: true` flag so the UI and DELETE
endpoint know to treat it specially. Idempotent; safe to call on
every boot. Doesn't overwrite the file once it exists (so user edits
are preserved across restarts and upgrades)."""
index = _load_index()
index = p_load_index()
dirty = False
for entry in _built_in_skill_registry():
for entry in p_built_in_skill_registry():
skill_id = entry["id"]
fpath = os.path.join(SKILLS_DIR, f"{skill_id}.md")
if not os.path.exists(fpath):
@@ -103,7 +103,7 @@ def _seed_built_in_skills() -> None:
index[skill_id] = meta
dirty = True
if dirty:
_save_index(index)
p_save_index(index)
@asynccontextmanager
@@ -111,7 +111,7 @@ async def skills_lifespan():
os.makedirs(SKILLS_DIR, exist_ok=True)
os.makedirs(SKILLS_WORKSPACE_DIR, exist_ok=True)
try:
_seed_built_in_skills()
p_seed_built_in_skills()
except Exception:
# Don't block app startup on a skill-seed failure; the worst
# case is the user has to manually paste the skill in once.
@@ -122,9 +122,9 @@ async def skills_lifespan():
skills = SubApp("skills", skills_lifespan)
def _sync_skills() -> list[Skill]:
def p_sync_skills() -> list[Skill]:
"""Sync skills from the filesystem, updating the index."""
index = _load_index()
index = p_load_index()
result = []
if os.path.exists(SKILLS_DIR):
@@ -152,10 +152,10 @@ def _sync_skills() -> list[Skill]:
@skills.router.get("/list")
async def list_skills():
return {"skills": [s.model_dump() for s in _sync_skills()]}
return {"skills": [s.model_dump() for s in p_sync_skills()]}
def _parse_skill_frontmatter(raw: str) -> dict:
def p_parse_skill_frontmatter(raw: str) -> dict:
"""Extract YAML frontmatter fields from a SKILL.md file."""
if not raw.startswith("---"):
return {}
@@ -207,7 +207,7 @@ async def read_skill_workspace(workspace_id: str):
except json.JSONDecodeError:
pass
frontmatter = _parse_skill_frontmatter(skill_content) if skill_content else {}
frontmatter = p_parse_skill_frontmatter(skill_content) if skill_content else {}
return {
"skill_content": skill_content,
@@ -218,7 +218,7 @@ async def read_skill_workspace(workspace_id: str):
@skills.router.get("/{skill_id}")
async def get_skill(skill_id: str):
for s in _sync_skills():
for s in p_sync_skills():
if s.id == skill_id:
return s.model_dump()
raise HTTPException(status_code=404, detail="Skill not found")
@@ -232,13 +232,13 @@ async def create_skill(body: SkillCreate):
with open(fpath, "w") as f:
f.write(body.content)
index = _load_index()
index = p_load_index()
index[slug] = {
"name": body.name,
"description": body.description,
"command": body.command or slug,
}
_save_index(index)
p_save_index(index)
skill = Skill(
id=slug,
@@ -262,7 +262,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
with open(fpath, "w") as f:
f.write(body.content)
index = _load_index()
index = p_load_index()
meta = index.get(skill_id, {})
if body.name is not None:
meta["name"] = body.name
@@ -271,7 +271,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
if body.command is not None:
meta["command"] = body.command
index[skill_id] = meta
_save_index(index)
p_save_index(index)
with open(fpath) as f:
content = f.read()
@@ -289,7 +289,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
@skills.router.delete("/{skill_id}")
async def delete_skill(skill_id: str):
index = _load_index()
index = p_load_index()
if index.get(skill_id, {}).get("built_in"):
raise HTTPException(
status_code=409,
@@ -303,5 +303,5 @@ async def delete_skill(skill_id: str):
if os.path.exists(fpath):
os.remove(fpath)
index.pop(skill_id, None)
_save_index(index)
p_save_index(index)
return {"ok": True}
+20 -16
View File
@@ -26,10 +26,10 @@ logger = logging.getLogger(__name__)
# Namespaces the hash so a raw hardware UUID never leaves the device. Public on
# purpose (open-source): it only prevents transmitting the raw id, not a secret.
_FP_SALT = "openswarm-free-trial-v1"
P_FP_SALT = "openswarm-free-trial-v1"
def _enabled() -> bool:
def p_enabled() -> bool:
# Default ON as of 1.2.80: the cloud free-trial proxy is live on prod
# (api.openswarm.com) and arming + metered Haiku were verified end to end.
# Set OPENSWARM_FREE_TRIAL_ENABLED=0 to force it off. The pool-shed gate +
@@ -38,7 +38,7 @@ def _enabled() -> bool:
return os.environ.get("OPENSWARM_FREE_TRIAL_ENABLED", "1") == "1"
def _raw_hardware_id() -> str | None:
def p_raw_hardware_id() -> str | None:
"""A stable per-machine id that survives app reinstall / data wipe."""
system = platform.system()
try:
@@ -67,17 +67,18 @@ def _raw_hardware_id() -> str | None:
return None
def _fingerprint(settings_obj) -> str | None:
raw = _raw_hardware_id()
def p_fingerprint(settings_obj) -> str | None:
raw = p_raw_hardware_id()
if not raw:
# Fail-soft: installation_id is less durable (regenerates on wipe) but
# better than nothing on a machine where the hardware id can't be read.
raw = getattr(settings_obj, "installation_id", None)
if not raw:
return None
return hashlib.sha256((_FP_SALT + raw).encode("utf-8")).hexdigest()
return hashlib.sha256((P_FP_SALT + raw).encode("utf-8")).hexdigest()
# Public - called by settings.py
def has_own_model(s) -> bool:
"""True if the user already has any real model path in settings; never shadow it."""
if any(getattr(s, k, None) for k in (
@@ -95,7 +96,7 @@ def has_own_model(s) -> bool:
return False
async def _has_connected_subscription() -> bool:
async def p_has_connected_subscription() -> bool:
"""True if 9Router holds a live Claude/ChatGPT/Gemini subscription. Those
connections live in 9Router, not settings, so the sync check above misses
them; this catches a sub connected while the trial was armed."""
@@ -112,11 +113,11 @@ async def _has_connected_subscription() -> bool:
return False
def _proxy_base(settings_obj) -> str:
def p_proxy_base(settings_obj) -> str:
return (getattr(settings_obj, "openswarm_proxy_url", None) or OPENSWARM_DEFAULT_PROXY_URL).rstrip("/")
async def _sync_routing(settings_obj) -> None:
async def p_sync_routing(settings_obj) -> None:
try:
from backend.apps.nine_router.sync_custom import sync_pro_routing
await sync_pro_routing(settings_obj)
@@ -124,6 +125,7 @@ async def _sync_routing(settings_obj) -> None:
logger.debug("free-trial routing sync skipped: %s", e)
# Public - called by agent_manager.py
async def clear_free_trial(settings_obj) -> None:
"""Drop the trial token and revert to own_key. Keeps free_trial_remaining
(so the UI knows it's spent) and never touches a real paid mode."""
@@ -131,29 +133,30 @@ async def clear_free_trial(settings_obj) -> None:
settings_obj.connection_mode = "own_key"
settings_obj.free_trial_token = None
await save_settings_async(settings_obj)
await _sync_routing(settings_obj)
await p_sync_routing(settings_obj)
# Public - called by subscription.router.py
async def arm_free_trial(settings_obj) -> dict:
"""Mint (or re-fetch) the machine's grant and, if runs remain, flip into
free-trial mode. Guarded: never arms over a real key/subscription."""
if not _enabled():
if not p_enabled():
return {"armed": False, "reason": "disabled"}
mode = getattr(settings_obj, "connection_mode", "own_key")
if mode not in ("own_key", "free-trial"):
return {"armed": False, "reason": "other_mode"}
if has_own_model(settings_obj) or await _has_connected_subscription():
if has_own_model(settings_obj) or await p_has_connected_subscription():
# A real model exists now (key, custom provider, or a 9Router sub). If we
# were on the free lane, hand the wheel back instead of re-arming.
if mode == "free-trial":
await clear_free_trial(settings_obj)
return {"armed": False, "reason": "has_model"}
fp = _fingerprint(settings_obj)
fp = p_fingerprint(settings_obj)
if not fp:
return {"armed": False, "reason": "no_fingerprint"}
base = _proxy_base(settings_obj)
base = p_proxy_base(settings_obj)
payload: dict = {"fingerprint_hash": fp}
if getattr(settings_obj, "installation_id", None):
payload["install_id"] = settings_obj.installation_id
@@ -185,7 +188,7 @@ async def arm_free_trial(settings_obj) -> dict:
if not cur.startswith(("sonnet", "haiku", "opus")):
settings_obj.default_model = "haiku"
await save_settings_async(settings_obj)
await _sync_routing(settings_obj)
await p_sync_routing(settings_obj)
return {"armed": True, "runs_remaining": remaining, "runs_limit": settings_obj.free_trial_runs_limit}
# Already spent on this machine: record it but don't arm.
@@ -193,6 +196,7 @@ async def arm_free_trial(settings_obj) -> dict:
return {"armed": False, "reason": "exhausted", "runs_remaining": 0}
# Public - called by subscription.router.py
async def refresh_free_trial(settings_obj) -> dict:
"""Re-read remaining runs from the cloud. Called after a session ends so the
onboarding 'runs low' nudge stays honest. Clears the trial when spent."""
@@ -200,7 +204,7 @@ async def refresh_free_trial(settings_obj) -> dict:
if getattr(settings_obj, "connection_mode", "own_key") != "free-trial" or not token:
return {"connected": False, "runs_remaining": getattr(settings_obj, "free_trial_remaining", None)}
base = _proxy_base(settings_obj)
base = p_proxy_base(settings_obj)
try:
async with httpx.AsyncClient(timeout=5.0) as client:
r = await client.post(
+24 -24
View File
@@ -26,7 +26,7 @@ async def subscription_lifespan():
subscription = SubApp("subscription", subscription_lifespan)
def _proxy_url() -> str:
def p_proxy_url() -> str:
"""Cloud router base URL. Overridable per-user via settings, falling back
to the module-default. No trailing slash."""
settings_obj = load_settings()
@@ -35,7 +35,7 @@ def _proxy_url() -> str:
return url.rstrip("/")
async def _sync_pro_routing(settings_obj) -> None:
async def p_sync_pro_routing(settings_obj) -> None:
"""Mirror connection state into 9Router's Claude lane (WebSearch on
non-Claude primaries). PUT /api/settings no longer carries these fields,
so the state-change endpoints here are the only trigger left."""
@@ -46,7 +46,7 @@ async def _sync_pro_routing(settings_obj) -> None:
logger.debug("pro routing sync skipped: %s", e)
async def _clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None:
async def p_clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None:
"""Revert to own_key mode and drop OpenSwarm Pro routing state.
`drop_bearer=True` (the default) is the original behavior, used when the
@@ -66,17 +66,17 @@ async def _clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None
settings_obj.openswarm_subscription_expires = None
settings_obj.openswarm_usage_cached = None
await save_settings_async(settings_obj)
_sync_subscription_identity(settings_obj)
await _sync_pro_routing(settings_obj)
p_sync_subscription_identity(settings_obj)
await p_sync_pro_routing(settings_obj)
def _sync_subscription_identity(settings_obj) -> None:
def p_sync_subscription_identity(settings_obj) -> None:
"""Push the installation's current subscription state into service-sync person
properties so every event from this user is segmentable by plan /
paying-vs-free. Safe to call from hot paths; service-sync is fire-and-forget
and swallows errors internally."""
try:
from backend.apps.service.client import identify as _identify
from backend.apps.service.client import identify
except Exception:
return
mode = getattr(settings_obj, "connection_mode", "own_key")
@@ -92,7 +92,7 @@ def _sync_subscription_identity(settings_obj) -> None:
if is_paying and expires:
props["subscription_expires"] = expires
try:
_identify(props)
identify(props)
except Exception as e:
logger.debug("identify sync failed: %s", e)
@@ -118,7 +118,7 @@ async def activate(body: ActivateRequest):
if not body.token or len(body.token) < 16:
raise HTTPException(status_code=400, detail="Invalid token")
proxy = _proxy_url()
proxy = p_proxy_url()
try:
async with httpx.AsyncClient(timeout=10.0) as client:
r = await client.get(
@@ -165,8 +165,8 @@ async def activate(body: ActivateRequest):
settings_obj.openswarm_usage_cached = usage
await save_settings_async(settings_obj)
_sync_subscription_identity(settings_obj)
await _sync_pro_routing(settings_obj)
p_sync_subscription_identity(settings_obj)
await p_sync_pro_routing(settings_obj)
return {"ok": True, "plan": settings_obj.openswarm_subscription_plan}
@@ -199,7 +199,7 @@ async def status():
try:
async with httpx.AsyncClient(timeout=5.0) as client:
r = await client.get(
f"{_proxy_url()}/api/me",
f"{p_proxy_url()}/api/me",
headers={"Authorization": f"Bearer {bearer}"},
)
upstream_code = r.status_code
@@ -219,7 +219,7 @@ async def status():
# through a dead subscription. Settings UI sees connected=False and
# falls back to the Subscribe CTA; chat reverts to own_key routing.
if upstream_code in (401, 402):
await _clear_subscription(settings_obj)
await p_clear_subscription(settings_obj)
return {
"connected": False,
"connection_mode": "own_key",
@@ -253,34 +253,34 @@ async def sync():
already had."""
# Lazy-import the service-sync helper so subscription/router doesn't pay the
# cost when analytics are disabled.
from backend.apps.service.client import sync as _sync
from backend.apps.service.client import sync as client_sync
settings_obj = load_settings()
bearer = getattr(settings_obj, "openswarm_bearer_token", None)
mode = getattr(settings_obj, "connection_mode", "own_key")
if mode != "openswarm-pro" or not bearer:
_sync(settings_obj.model_dump())
client_sync(settings_obj.model_dump())
return {"ok": True, "synced": False, "connection_mode": mode}
try:
async with httpx.AsyncClient(timeout=10.0) as client:
r = await client.post(
f"{_proxy_url()}/api/subscription/sync",
f"{p_proxy_url()}/api/subscription/sync",
headers={"Authorization": f"Bearer {bearer}"},
)
except httpx.HTTPError as e:
logger.debug("subscription/sync live fetch failed: %s", e)
_sync(settings_obj.model_dump())
client_sync(settings_obj.model_dump())
return {"ok": True, "synced": False, "reason": "network"}
# Same 401/402 handling as /status: if Stripe-side reconciliation proves
# the bearer is dead or the sub expired, clear local state so the app
# reverts to own_key instead of hammering a useless token.
if r.status_code in (401, 402):
await _clear_subscription(settings_obj)
await p_clear_subscription(settings_obj)
reason = "revoked" if r.status_code == 401 else "expired"
_sync(settings_obj.model_dump())
client_sync(settings_obj.model_dump())
return {
"ok": True,
"synced": False,
@@ -290,7 +290,7 @@ async def sync():
if r.status_code != 200:
logger.debug("subscription/sync got %s from cloud: %s", r.status_code, r.text[:200])
_sync(settings_obj.model_dump())
client_sync(settings_obj.model_dump())
return {"ok": True, "synced": False, "reason": "upstream"}
data = r.json()
@@ -307,8 +307,8 @@ async def sync():
datetime.fromtimestamp(period_end_ms / 1000, tz=timezone.utc).isoformat()
)
await save_settings_async(settings_obj)
_sync_subscription_identity(settings_obj)
_sync(settings_obj.model_dump())
p_sync_subscription_identity(settings_obj)
client_sync(settings_obj.model_dump())
return {
"ok": True,
"synced": bool(data.get("synced")),
@@ -333,7 +333,7 @@ async def portal():
async with httpx.AsyncClient(timeout=10.0) as client:
r = await client.post(
f"{_proxy_url()}/api/billing/portal",
f"{p_proxy_url()}/api/billing/portal",
headers={"Authorization": f"Bearer {bearer}"},
)
if r.status_code >= 400:
@@ -373,5 +373,5 @@ async def disconnect():
does NOT sign the user out of OpenSwarm (use /api/auth/signout for that).
Useful when a user wants to temporarily route through their own API key
without losing their account state."""
await _clear_subscription(load_settings(), drop_bearer=False)
await p_clear_subscription(load_settings(), drop_bearer=False)
return {"ok": True}
+42 -38
View File
@@ -11,18 +11,19 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
logger = logging.getLogger(__name__)
# Public - called by main.py, agent_manager.py, prompt_context.py
def sanitize_mcp_server_name(name: str) -> str:
"""Convert a tool name into a valid MCP server identifier (alphanumeric + hyphens)."""
return re.sub(r"[^a-z0-9]+", "-", name.lower()).strip("-")
def _extra_bin_dirs() -> list[str]:
def p_extra_bin_dirs() -> list[str]:
"""Well-known user-local bin directories that may not be on PATH in packaged apps."""
home = os.path.expanduser("~")
# Bundled uv-bin (ships uvx for non-dev users)
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
dirs = [
os.path.join(_backend, "uv-bin"),
os.path.join(backend_dir, "uv-bin"),
os.path.join(home, ".bun", "bin"),
os.path.join(home, ".cargo", "bin"),
os.path.join(home, ".local", "bin"),
@@ -46,7 +47,8 @@ def _extra_bin_dirs() -> list[str]:
return dirs
def _resolve_command(command: str) -> str | None:
# Public - called by mcp_discovery.py
def resolve_command(command: str) -> str | None:
"""Find a command on PATH, falling back to common user-local bin directories
and bundled binaries (uv-bin for uvx/uv)."""
found = shutil.which(command)
@@ -59,24 +61,25 @@ def _resolve_command(command: str) -> str | None:
suffixes = [""] + os.environ.get("PATHEXT", ".COM;.EXE;.BAT;.CMD").lower().split(os.pathsep)
else:
suffixes = [""]
def _probe(directory: str) -> str | None:
def probe(directory: str) -> str | None:
for suffix in suffixes:
candidate = os.path.join(directory, command + suffix)
if os.path.isfile(candidate) and os.access(candidate, os.X_OK):
return candidate
return None
for d in _extra_bin_dirs():
hit = _probe(d)
for d in p_extra_bin_dirs():
hit = probe(d)
if hit:
return hit
# Check bundled uv-bin directory (ships uv/uvx for non-dev users)
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
return _probe(os.path.join(_backend, "uv-bin"))
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
return probe(os.path.join(backend_dir, "uv-bin"))
def _augmented_path() -> str:
# Public - called by mcp_discovery.py
def augmented_path() -> str:
"""Return PATH with extra bin dirs prepended (for child process environments)."""
extra = [d for d in _extra_bin_dirs() if os.path.isdir(d)]
extra = [d for d in p_extra_bin_dirs() if os.path.isdir(d)]
current = os.environ.get("PATH", "")
seen: set[str] = set()
parts: list[str] = []
@@ -87,6 +90,7 @@ def _augmented_path() -> str:
return os.pathsep.join(parts)
# Public - called by tools_lib.py
def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
"""Build the claude_agent_sdk mcp_servers config entry for a tool.
@@ -130,9 +134,9 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
# proxy that forwards the refresh to our cloud's pool-aware
# /api/oauth/google/refresh endpoint; CLIENT_ID/SECRET become
# unused placeholders (gauth.py only validates non-empty).
_port = os.environ.get("OPENSWARM_PORT", "8324")
port = os.environ.get("OPENSWARM_PORT", "8324")
env["GOOGLE_WORKSPACE_TOKEN_URI"] = (
f"http://127.0.0.1:{_port}/api/tools/google-oauth-token"
f"http://127.0.0.1:{port}/api/tools/google-oauth-token"
)
env.setdefault("GOOGLE_WORKSPACE_CLIENT_ID", "openswarm-proxy")
env.setdefault("GOOGLE_WORKSPACE_CLIENT_SECRET", "openswarm-proxy")
@@ -167,9 +171,9 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
# The shim runs as a subprocess and needs to import
# `backend.apps.discord_mcp_shim`; set PYTHONPATH to the project
# root (parent of the backend/ dir) so that import resolves.
_project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
existing_pp = env.get("PYTHONPATH") or os.environ.get("PYTHONPATH", "")
env["PYTHONPATH"] = (_project_root + os.pathsep + existing_pp) if existing_pp else _project_root
env["PYTHONPATH"] = (project_root + os.pathsep + existing_pp) if existing_pp else project_root
# Microsoft 365 MCP: use a stable token cache path shared across process spawns
if tool.name.lower() == "microsoft 365" and config.get("type") == "stdio":
@@ -196,7 +200,7 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
if config["command"] in ("npx", "bunx"):
pkg_name = next((a for a in (config.get("args") or []) if not a.startswith("-")), None)
if pkg_name:
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
electron_path = os.environ.get("OPENSWARM_ELECTRON_PATH")
# Two bundle layouts in mcp-bundles/, checked in priority order:
#
@@ -217,8 +221,8 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
# Scoped names get flattened ("@softeria/ms-365-mcp-server"
# -> "softeria-ms-365-mcp-server") for filesystem safety.
safe_bundle = pkg_name.replace("/", "-").replace("@", "")
bundle_dir_path = os.path.join(_backend, "mcp-bundles", safe_bundle, "dist", "index.js")
bundle_file_path = os.path.join(_backend, "mcp-bundles", f"{safe_bundle}.js")
bundle_dir_path = os.path.join(backend_dir, "mcp-bundles", safe_bundle, "dist", "index.js")
bundle_file_path = os.path.join(backend_dir, "mcp-bundles", f"{safe_bundle}.js")
bundle_path = None
if os.path.isfile(bundle_dir_path):
bundle_path = bundle_dir_path
@@ -242,12 +246,12 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
else:
# Check for pre-installed npm package (works in both dev and packaged modes)
safe_dir = pkg_name.replace("/", "-").replace("@", "")
npm_dir = os.path.join(_backend, "npm-servers", safe_dir)
npm_dir = os.path.join(backend_dir, "npm-servers", safe_dir)
pkg_json_path = os.path.join(npm_dir, "node_modules", pkg_name, "package.json")
if os.path.isfile(pkg_json_path):
import json as _json
import json
with open(pkg_json_path) as f:
pkg_meta = _json.load(f)
pkg_meta = json.load(f)
bin_field = pkg_meta.get("bin", {})
entry = list(bin_field.values())[0] if isinstance(bin_field, dict) else bin_field
# Same priority as 9Router / MCP-bundle paths: bundled node > system node > Electron-as-Node.
@@ -262,33 +266,33 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
logger.info(f"Using pre-installed npm MCP server for {pkg_name}")
if not os.path.isabs(config.get("command", "")):
resolved = _resolve_command(config["command"])
resolved = resolve_command(config["command"])
if resolved:
config["command"] = resolved
else:
logger.warning(f"Command '{config['command']}' not found on PATH or bundled directories")
env = config.setdefault("env", {})
env.setdefault("PATH", _augmented_path())
env.setdefault("PATH", augmented_path())
env.setdefault("PYTHONPATH", "")
# Point uv/uvx at our bundled Python; avoids macOS CLT popup on fresh Macs
# and avoids downloading Python at runtime
_is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
_is_windows = sys.platform == "win32"
if _is_packaged:
_resources = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
if _is_windows:
_bundled_python = os.path.join(_resources, "python-env", "python.exe")
is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
is_windows = sys.platform == "win32"
if is_packaged:
resources_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
if is_windows:
bundled_python = os.path.join(resources_dir, "python-env", "python.exe")
else:
_bundled_python = os.path.join(_resources, "python-env", "bin", "python3")
if os.path.exists(_bundled_python):
env.setdefault("UV_PYTHON", _bundled_python)
bundled_python = os.path.join(resources_dir, "python-env", "bin", "python3")
if os.path.exists(bundled_python):
env.setdefault("UV_PYTHON", bundled_python)
else:
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
if _is_windows:
_venv_python = os.path.join(_backend, ".venv", "Scripts", "python.exe")
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
if is_windows:
venv_python = os.path.join(backend_dir, ".venv", "Scripts", "python.exe")
else:
_venv_python = os.path.join(_backend, ".venv", "bin", "python3")
if os.path.exists(_venv_python):
env.setdefault("UV_PYTHON", _venv_python)
venv_python = os.path.join(backend_dir, ".venv", "bin", "python3")
if os.path.exists(venv_python):
env.setdefault("UV_PYTHON", venv_python)
return config
+26 -23
View File
@@ -8,12 +8,12 @@ import shutil
import httpx
from fastapi import HTTPException
from backend.apps.tools_lib.mcp_config import _augmented_path, _resolve_command
from backend.apps.tools_lib.mcp_config import augmented_path, resolve_command
logger = logging.getLogger(__name__)
def _parse_sse_json(text: str) -> dict | None:
def p_parse_sse_json(text: str) -> dict | None:
"""Extract JSON from an SSE response body (handles `data: {...}` lines)."""
for line in text.splitlines():
stripped = line.strip()
@@ -30,7 +30,8 @@ def _parse_sse_json(text: str) -> dict | None:
return None
async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> list[dict]:
# Public - called by tools_lib.py
async def discover_mcp_tools_http(url: str, headers: dict | None = None) -> list[dict]:
"""Connect to a Streamable HTTP MCP server and call tools/list via JSON-RPC POST."""
h = {
"Content-Type": "application/json",
@@ -62,7 +63,7 @@ async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> lis
ct = list_resp.headers.get("content-type", "")
if "text/event-stream" in ct:
data = _parse_sse_json(list_resp.text)
data = p_parse_sse_json(list_resp.text)
else:
data = list_resp.json()
@@ -73,7 +74,8 @@ async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> lis
return [{"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} for t in tools_list]
async def _discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list[dict]:
# Public - called by tools_lib.py
async def discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list[dict]:
"""Connect to a legacy SSE MCP server (GET event-stream + POST messages) and call tools/list."""
from mcp.client.sse import sse_client
from mcp import ClientSession
@@ -99,10 +101,10 @@ async def _discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list
raise HTTPException(status_code=502, detail=f"SSE discovery failed: {first}") from first
_NPX_CACHE_RE = re.compile(r"_npx[/\\]([0-9a-f]{8,})[/\\]")
P_NPX_CACHE_RE = re.compile(r"_npx[/\\]([0-9a-f]{8,})[/\\]")
def _try_heal_npx_cache(stderr: str) -> str | None:
def p_try_heal_npx_cache(stderr: str) -> str | None:
"""On `ERR_MODULE_NOT_FOUND` pointing into `~/.npm/_npx/<hash>/`, wipe that one dir.
Why: interrupted npx installs leave a `package-lock.json` in the cache dir so
@@ -112,7 +114,7 @@ def _try_heal_npx_cache(stderr: str) -> str | None:
"""
if "ERR_MODULE_NOT_FOUND" not in stderr:
return None
m = _NPX_CACHE_RE.search(stderr)
m = P_NPX_CACHE_RE.search(stderr)
if not m:
return None
hash_ = m.group(1)
@@ -124,18 +126,19 @@ def _try_heal_npx_cache(stderr: str) -> str | None:
return hash_
async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None, env: dict | None = None, _attempt: int = 0) -> list[dict]:
# Public - called by tools_lib.py
async def discover_mcp_tools_stdio(command: str, args: list[str] | None = None, env: dict | None = None, attempt: int = 0) -> list[dict]:
"""Spawn a stdio MCP server process and call tools/list via JSON-RPC over stdin/stdout.
On the first attempt, a failure that looks like corrupted npx cache
(`ERR_MODULE_NOT_FOUND` pointing into `~/.npm/_npx/<hash>/`) triggers one
auto-heal + retry. No heal on `_attempt >= 1`.
auto-heal + retry. No heal on `attempt >= 1`.
"""
cmd_path = _resolve_command(command)
cmd_path = resolve_command(command)
if not cmd_path:
raise HTTPException(status_code=400, detail=f"Command '{command}' not found on PATH or common install locations")
proc_env = {**os.environ, **(env or {}), "PATH": _augmented_path()}
proc_env = {**os.environ, **(env or {}), "PATH": augmented_path()}
proc_env.pop("PYTHONPATH", None)
proc = await asyncio.create_subprocess_exec(
@@ -155,7 +158,7 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
# the opaque "discovery failed" we used to show.
stderr_tail: list[str] = []
async def _drain_stderr() -> None:
async def drain_stderr() -> None:
try:
while True:
chunk = await proc.stderr.readline()
@@ -169,14 +172,14 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
except Exception:
return
stderr_task = asyncio.create_task(_drain_stderr())
stderr_task = asyncio.create_task(drain_stderr())
async def _send(msg: dict) -> None:
async def send(msg: dict) -> None:
line = json.dumps(msg) + "\n"
proc.stdin.write(line.encode())
await proc.stdin.drain()
async def _recv(timeout_s: float = 30.0) -> dict:
async def recv(timeout_s: float = 30.0) -> dict:
"""Read JSON-RPC responses, skipping notification lines (no 'id' field)."""
while True:
line = await asyncio.wait_for(proc.stdout.readline(), timeout=timeout_s)
@@ -204,7 +207,7 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
return data
try:
await _send({
await send({
"jsonrpc": "2.0", "id": 1, "method": "initialize",
"params": {
"protocolVersion": "2025-03-26",
@@ -217,12 +220,12 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
# AV-scan every file npm writes; total install time often exceeds
# 60 s and occasionally pushes past 90 s. Subsequent reads run
# against an already-running server and stay at the default 30 s.
await _recv(timeout_s=120.0)
await recv(timeout_s=120.0)
await _send({"jsonrpc": "2.0", "method": "notifications/initialized"})
await send({"jsonrpc": "2.0", "method": "notifications/initialized"})
await _send({"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}})
data = await _recv()
await send({"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}})
data = await recv()
tools_list = data.get("result", {}).get("tools", [])
return [{"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} for t in tools_list]
@@ -231,8 +234,8 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
# Heal-on-corrupt-npx-cache still triggers from the EOF branch,
# which now includes the full stderr tail in `e.detail`; so the
# ERR_MODULE_NOT_FOUND signature is still discoverable here.
if _attempt == 0 and _try_heal_npx_cache(str(e.detail) if e.detail is not None else ""):
return await _discover_mcp_tools_stdio(command, args, env, _attempt=1)
if attempt == 0 and p_try_heal_npx_cache(str(e.detail) if e.detail is not None else ""):
return await discover_mcp_tools_stdio(command, args, env, attempt=1)
raise
except asyncio.TimeoutError:
# Most common cause: cold npx cache on Windows. The npm install
+23 -17
View File
@@ -13,7 +13,7 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
logger = logging.getLogger(__name__)
def _save(tool: ToolDefinition) -> None:
def p_save(tool: ToolDefinition) -> None:
with open(os.path.join(DATA_DIR, f"{tool.id}.json"), "w") as f:
json.dump(tool.model_dump(), f, indent=2)
@@ -22,7 +22,7 @@ def _save(tool: ToolDefinition) -> None:
# through the Fly cloud-proxy so client_secret values never ship inside the
# desktop binary. v1.0.28 was the last release that used a local Google
# callback with the client_secret in backend/.env.
_TOOL_NAME_TO_PROVIDER = {
P_TOOL_NAME_TO_PROVIDER = {
"airtable": "airtable",
"hubspot": "hubspot",
"discord": "discord",
@@ -34,11 +34,13 @@ _TOOL_NAME_TO_PROVIDER = {
}
def _proxied_provider_for(tool: ToolDefinition) -> Optional[str]:
return _TOOL_NAME_TO_PROVIDER.get(tool.name.lower())
# Public - called by tools_lib.py
def proxied_provider_for(tool: ToolDefinition) -> Optional[str]:
return P_TOOL_NAME_TO_PROVIDER.get(tool.name.lower())
def _persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
# Public - called by tools_lib.py
def persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
"""Normalise the cloud's claim response into tool.oauth_tokens.
Per-provider shaping mirrors what the v1.0.25 local-callback flow used
@@ -77,7 +79,7 @@ def _persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
tool.auth_status = "connected"
async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry: int) -> Optional[str]:
async def p_refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry: int) -> Optional[str]:
"""Refresh an OAuth access_token by POSTing the refresh_token to the
helper service. Per-provider wrappers below pass a default expires_in
fallback for providers that don't return one.
@@ -101,7 +103,7 @@ async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry
# Provider rejected; user revoked at the provider's side. Mark
# as needing re-auth so the UI prompts a Reconnect.
tool.auth_status = "expired"
_save(tool)
p_save(tool)
logger.warning(f"{provider} refresh rejected (user revoked); marking tool as expired")
return None
if resp.status_code != 200:
@@ -121,13 +123,14 @@ async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry
# Backfill identity label on first successful refresh after upgrade.
if not tool.connected_account_email and data.get("email"):
tool.connected_account_email = data["email"]
_save(tool)
p_save(tool)
return new_token
except Exception as e:
logger.warning(f"{provider} cloud refresh exception for tool {tool.id}: {e}")
return None
# Public - called by tools_lib.py
async def refresh_google_token(tool: ToolDefinition) -> Optional[str]:
"""Refresh an expired Google access_token via the Fly cloud-proxy.
@@ -135,20 +138,22 @@ async def refresh_google_token(tool: ToolDefinition) -> Optional[str]:
refresh_token. Same pattern as Airtable/HubSpot. Pre-v1.0.29 builds
held the secret in their bundled .env; v1.0.29 removed it.
"""
return await _refresh_via_proxy("google", tool, default_expiry=3600)
return await p_refresh_via_proxy("google", tool, default_expiry=3600)
# Public - called by tools_lib.py
async def refresh_airtable_token(tool: ToolDefinition) -> Optional[str]:
"""Refresh an expired Airtable OAuth access_token."""
return await _refresh_via_proxy("airtable", tool, default_expiry=7200)
return await p_refresh_via_proxy("airtable", tool, default_expiry=7200)
# Public - called by tools_lib.py
async def refresh_hubspot_token(tool: ToolDefinition) -> Optional[str]:
"""Refresh an expired HubSpot OAuth access_token."""
return await _refresh_via_proxy("hubspot", tool, default_expiry=1800)
return await p_refresh_via_proxy("hubspot", tool, default_expiry=1800)
def _m365_server_script() -> str:
# Public - called by tools_lib.py
def m365_server_script() -> str:
"""Return the on-disk path to the bundled MS365 MCP server entry.
v1.0.26 replaced the heavy backend/npm-servers/softeria-ms-365-mcp-server/
@@ -158,9 +163,9 @@ def _m365_server_script() -> str:
package.json) because cli.js reads __dirname/../package.json for the
--version flag; see scripts/build-app.sh `build_mcp_bundle_dir`.
"""
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
p_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
bundle = os.path.join(
_backend, "mcp-bundles", "softeria-ms-365-mcp-server", "dist", "index.js",
p_backend, "mcp-bundles", "softeria-ms-365-mcp-server", "dist", "index.js",
)
if os.path.isfile(bundle):
return bundle
@@ -168,12 +173,13 @@ def _m365_server_script() -> str:
# was left over from before the bundle migration. Will return the legacy
# path; if that doesn't exist either, the caller raises a clear error.
return os.path.join(
_backend, "npm-servers", "softeria-ms-365-mcp-server",
p_backend, "npm-servers", "softeria-ms-365-mcp-server",
"node_modules", "@softeria", "ms-365-mcp-server", "dist", "index.js",
)
def _m365_cache_env() -> dict[str, str]:
# Public - called by tools_lib.py
def m365_cache_env() -> dict[str, str]:
cache_dir = os.path.join(os.path.expanduser("~"), ".openswarm")
os.makedirs(cache_dir, exist_ok=True)
return {
+15 -14
View File
@@ -1,8 +1,8 @@
_READ_PREFIXES = ("get", "list", "read", "search", "fetch", "find", "query", "count", "check", "describe", "show", "download", "browse", "analy", "explain")
_WRITE_PREFIXES = ("create", "write", "delete", "update", "send", "remove", "modify", "add", "set", "put", "post", "patch", "insert", "move", "copy", "rename", "archive", "trash", "publish", "approve", "reject")
P_READ_PREFIXES = ("get", "list", "read", "search", "fetch", "find", "query", "count", "check", "describe", "show", "download", "browse", "analy", "explain")
P_WRITE_PREFIXES = ("create", "write", "delete", "update", "send", "remove", "modify", "add", "set", "put", "post", "patch", "insert", "move", "copy", "rename", "archive", "trash", "publish", "approve", "reject")
_SERVICE_RULES: list[tuple[list[str], str, str]] = [
P_SERVICE_RULES: list[tuple[list[str], str, str]] = [
# (keywords, service_name, group)
# Google Workspace
(["gmail"], "Gmail", "Google"),
@@ -31,20 +31,20 @@ _SERVICE_RULES: list[tuple[list[str], str, str]] = [
]
def _categorize_tool(name: str) -> str:
def p_categorize_tool(name: str) -> str:
lower = name.lower().replace("_", " ").replace("-", " ").strip()
for word in lower.split():
for prefix in _READ_PREFIXES:
for prefix in P_READ_PREFIXES:
if word.startswith(prefix):
return "read"
for prefix in _WRITE_PREFIXES:
for prefix in P_WRITE_PREFIXES:
if word.startswith(prefix):
return "write"
return "write"
def _integration_domain(integration: str) -> str:
"""Which curated _SERVICE_RULES set applies to this integration, if any. The Google rules use
def p_integration_domain(integration: str) -> str:
"""Which curated P_SERVICE_RULES set applies to this integration, if any. The Google rules use
generic words (message/table/page/doc/script) that otherwise mis-tag Slack/Notion/Airtable/M365."""
n = (integration or "").lower()
if "google" in n:
@@ -56,13 +56,13 @@ def _integration_domain(integration: str) -> str:
return ""
def _extract_service(name: str, integration: str) -> tuple[str, str]:
def p_extract_service(name: str, integration: str) -> tuple[str, str]:
"""Map a tool name to (service, group). Curated rulesets apply only to the integration they were
written for; every other integration groups under its own name so it isn't mislabeled as Google."""
domain = _integration_domain(integration)
domain = p_integration_domain(integration)
if domain:
lower = name.lower()
for keywords, display, group in _SERVICE_RULES:
for keywords, display, group in P_SERVICE_RULES:
if group != domain:
continue
for kw in keywords:
@@ -73,15 +73,16 @@ def _extract_service(name: str, integration: str) -> tuple[str, str]:
return (integration or "Other"), ""
def _classify_services(
# Public - called by tools_lib.py
def classify_services(
tool_names: list[str], integration: str
) -> tuple[dict[str, dict[str, list[str]]], dict[str, list[str]], list[str], list[str]]:
"""Bucket tool names into services + service groups + read/write categories for one integration."""
services: dict[str, dict[str, list[str]]] = {}
service_groups: dict[str, list[str]] = {}
for name in tool_names:
cat = _categorize_tool(name)
svc, group = _extract_service(name, integration)
cat = p_categorize_tool(name)
svc, group = p_extract_service(name, integration)
services.setdefault(svc, {"read": [], "write": []})
services[svc][cat].append(name)
if group:
+81 -77
View File
@@ -21,20 +21,20 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
# derive_mcp_config re-exported for agent_manager/main.
from backend.apps.tools_lib.mcp_config import derive_mcp_config
from backend.apps.tools_lib.mcp_discovery import (
_discover_mcp_tools_http,
_discover_mcp_tools_sse,
_discover_mcp_tools_stdio,
discover_mcp_tools_http,
discover_mcp_tools_sse,
discover_mcp_tools_stdio,
)
from backend.apps.tools_lib.tool_taxonomy import _classify_services
from backend.apps.tools_lib.tool_taxonomy import classify_services
# refresh_* re-exported for agent_manager.
from backend.apps.tools_lib.oauth_tokens import (
_proxied_provider_for,
_persist_cloud_tokens,
proxied_provider_for,
persist_cloud_tokens,
refresh_google_token,
refresh_airtable_token,
refresh_hubspot_token,
_m365_server_script,
_m365_cache_env,
m365_server_script,
m365_cache_env,
)
logger = logging.getLogger(__name__)
@@ -43,8 +43,8 @@ logger = logging.getLogger(__name__)
@asynccontextmanager
async def tools_lib_lifespan():
os.makedirs(DATA_DIR, exist_ok=True)
_ensure_default_permissions()
_reclassify_existing_tools()
p_ensure_default_permissions()
p_reclassify_existing_tools()
yield
@@ -55,10 +55,10 @@ tools_lib = SubApp("tools", tools_lib_lifespan)
# outputs (Gmail, WebFetch); every other built-in is sandboxed by domain.
# Must match agent_manager._DEFAULTS so the Settings UI and the agent agree
# on what "no policy set" means.
_DEFAULT_BUILTIN_POLICIES = {"Bash": "ask"}
P_DEFAULT_BUILTIN_POLICIES = {"Bash": "ask"}
def _ensure_default_permissions() -> None:
def p_ensure_default_permissions() -> None:
"""Seed BUILTIN_PERMISSIONS_PATH so the user's Settings toggles persist
cleanly. Without this the file is missing on first run, load returns {},
every PUT-from-the-UI overwrites with the partial payload the click
@@ -68,15 +68,15 @@ def _ensure_default_permissions() -> None:
"""
existing = load_builtin_permissions()
desired = {
t.name: _DEFAULT_BUILTIN_POLICIES.get(t.name, "always_allow")
t.name: P_DEFAULT_BUILTIN_POLICIES.get(t.name, "always_allow")
for t in BUILTIN_TOOLS
}
merged = {**desired, **existing}
merged = {**desired, **existing}
if merged != existing:
save_builtin_permissions(merged)
p_save_builtin_permissions(merged)
def _reclassify_existing_tools() -> None:
def p_reclassify_existing_tools() -> None:
"""One-time correction for tools discovered before service rules were integration-scoped: most
integrations got mislabeled under a bogus 'Google' group (generic keyword rules applied globally).
Recompute services/groups from each tool's stored tool names. Idempotent; rewrites only on change.
@@ -87,7 +87,7 @@ def _reclassify_existing_tools() -> None:
if not fname.endswith(".json"):
continue
try:
tool = _load(fname[:-5])
tool = p_load(fname[:-5])
except Exception:
continue
perms = tool.tool_permissions or {}
@@ -96,7 +96,7 @@ def _reclassify_existing_tools() -> None:
names = [k for k in perms if not k.startswith("_")]
if not names:
continue
services, service_groups, all_read, all_write = _classify_services(names, tool.name)
services, service_groups, all_read, all_write = classify_services(names, tool.name)
if perms.get("_services") == services and perms.get("_service_groups") == service_groups:
continue
perms["_services"] = services
@@ -104,7 +104,7 @@ def _reclassify_existing_tools() -> None:
perms["_categories"] = {"read": all_read, "write": all_write}
tool.tool_permissions = perms
try:
_save(tool)
p_save(tool)
except Exception:
pass
@@ -118,11 +118,11 @@ GOOGLE_USERINFO_URL = "https://www.googleapis.com/oauth2/v2/userinfo"
# MCPSearch keystroke; the cache skips re-parsing, revalidated by a per-file stat
# signature so any write (ours or external) invalidates instantly. Callers treat
# the returned ToolDefinitions as immutable; mutate via _load(tool_id) + _save.
_tools_cache: list[ToolDefinition] | None = None
_tools_cache_sig: tuple | None = None
P_TOOLS_CACHE: list[ToolDefinition] | None = None
P_TOOLS_CACHE_SIG: tuple | None = None
def _tools_sig() -> tuple | None:
def p_tools_sig() -> tuple | None:
if not os.path.exists(DATA_DIR):
return ()
try:
@@ -136,11 +136,12 @@ def _tools_sig() -> tuple | None:
return None
def _load_all() -> list[ToolDefinition]:
global _tools_cache, _tools_cache_sig
sig = _tools_sig()
if sig is not None and _tools_cache is not None and sig == _tools_cache_sig:
return list(_tools_cache)
# Public - called by mcp_pre_flight.py, prompt_context.py, main.py, agent_manager.py
def load_all_tools() -> list[ToolDefinition]:
global P_TOOLS_CACHE, P_TOOLS_CACHE_SIG
sig = p_tools_sig()
if sig is not None and P_TOOLS_CACHE is not None and sig == P_TOOLS_CACHE_SIG:
return list(P_TOOLS_CACHE)
result = []
if not os.path.exists(DATA_DIR):
return result
@@ -149,17 +150,17 @@ def _load_all() -> list[ToolDefinition]:
with open(os.path.join(DATA_DIR, fname)) as f:
result.append(ToolDefinition(**json.load(f)))
if sig is not None:
_tools_cache = list(result)
_tools_cache_sig = sig
P_TOOLS_CACHE = list(result)
P_TOOLS_CACHE_SIG = sig
return result
def _save(tool: ToolDefinition):
def p_save(tool: ToolDefinition):
with open(os.path.join(DATA_DIR, f"{tool.id}.json"), "w") as f:
json.dump(tool.model_dump(), f, indent=2)
def _load(tool_id: str) -> ToolDefinition:
def p_load(tool_id: str) -> ToolDefinition:
path = os.path.join(DATA_DIR, f"{tool_id}.json")
if not os.path.exists(path):
raise HTTPException(status_code=404, detail="Tool not found")
@@ -179,7 +180,7 @@ def _load(tool_id: str) -> ToolDefinition:
"command": "python",
"args": ["-m", "backend.apps.discord_mcp_shim"],
}
_save(tool)
p_save(tool)
return tool
@@ -188,6 +189,7 @@ async def list_builtin_tools():
return {"tools": [t.model_dump() for t in BUILTIN_TOOLS]}
# Public - called by agent_manager.py
def load_builtin_permissions() -> dict[str, str]:
if not os.path.exists(BUILTIN_PERMS_PATH):
return {}
@@ -195,12 +197,13 @@ def load_builtin_permissions() -> dict[str, str]:
return json.load(f)
def save_builtin_permissions(perms: dict[str, str]):
def p_save_builtin_permissions(perms: dict[str, str]):
os.makedirs(os.path.dirname(BUILTIN_PERMS_PATH), exist_ok=True)
with open(BUILTIN_PERMS_PATH, "w") as f:
json.dump(perms, f, indent=2)
# Public - called by agent_manager.py
def load_trusted_sensitive_paths() -> list[str]:
if not os.path.exists(TRUSTED_SENSITIVE_PATHS_PATH):
return []
@@ -215,6 +218,7 @@ def load_trusted_sensitive_paths() -> list[str]:
return [p for p in raw if isinstance(p, str) and p]
# Public - called by agent_manager.py
def save_trusted_sensitive_paths(patterns: list[str]):
os.makedirs(os.path.dirname(TRUSTED_SENSITIVE_PATHS_PATH), exist_ok=True)
seen: list[str] = []
@@ -254,14 +258,14 @@ async def update_builtin_permissions(body: dict):
for name, policy in body.get("permissions", {}).items():
if name in valid_tools and policy in valid_policies:
perms[name] = policy
save_builtin_permissions(perms)
p_save_builtin_permissions(perms)
return {"permissions": perms}
@tools_lib.router.get("/list")
async def list_tools():
tools = []
for t in _load_all():
for t in load_all_tools():
d = t.model_dump()
# Heal pre-fix tools whose persisted email is the "{name} account" placeholder so the pill stops reading like a name and falls back to plain "Connected".
placeholder = f"{t.name} account"
@@ -271,7 +275,7 @@ async def list_tools():
return {"tools": tools}
def _connected_html() -> HTMLResponse:
def p_connected_html() -> HTMLResponse:
"""v1.0.25-style auto-close page. Same markup so the UX is unchanged."""
return HTMLResponse("""
<html><body>
@@ -287,7 +291,7 @@ def _connected_html() -> HTMLResponse:
@tools_lib.router.get("/{tool_id}")
async def get_tool(tool_id: str):
return _load(tool_id).model_dump()
return p_load(tool_id).model_dump()
@tools_lib.router.post("/create")
@@ -301,16 +305,16 @@ async def create_tool(body: ToolCreate):
auth_type=body.auth_type,
auth_status=body.auth_status,
)
_save(tool)
p_save(tool)
return {"ok": True, "tool": tool.model_dump()}
@tools_lib.router.put("/{tool_id}")
async def update_tool(tool_id: str, body: ToolUpdate):
tool = _load(tool_id)
tool = p_load(tool_id)
for k, v in body.model_dump(exclude_none=True).items():
setattr(tool, k, v)
_save(tool)
p_save(tool)
return {"ok": True, "tool": tool.model_dump()}
@@ -324,7 +328,7 @@ async def delete_tool(tool_id: str):
@tools_lib.router.post("/{tool_id}/discover")
async def discover_tools(tool_id: str):
tool = _load(tool_id)
tool = p_load(tool_id)
if tool.auth_type == "oauth2" and tool.auth_status == "connected":
if tool.oauth_tokens.get("refresh_token"):
@@ -353,7 +357,7 @@ async def discover_tools(tool_id: str):
command = config.get("command", "")
if not command:
raise HTTPException(status_code=400, detail="stdio transport requires a 'command' in MCP config")
raw_tools = await _discover_mcp_tools_stdio(
raw_tools = await discover_mcp_tools_stdio(
command=command,
args=config.get("args"),
env=config.get("env"),
@@ -363,13 +367,13 @@ async def discover_tools(tool_id: str):
if not url:
raise HTTPException(status_code=400, detail="HTTP/SSE transport requires a 'url' in MCP config")
if transport == "sse":
raw_tools = await _discover_mcp_tools_sse(url, config.get("headers"))
raw_tools = await discover_mcp_tools_sse(url, config.get("headers"))
else:
try:
raw_tools = await _discover_mcp_tools_http(url, config.get("headers"))
raw_tools = await discover_mcp_tools_http(url, config.get("headers"))
except HTTPException:
logger.info(f"Streamable HTTP failed for {tool.name}, retrying with SSE transport")
raw_tools = await _discover_mcp_tools_sse(url, config.get("headers"))
raw_tools = await discover_mcp_tools_sse(url, config.get("headers"))
else:
raise HTTPException(status_code=400, detail=f"Unsupported MCP transport type: '{transport}'. Use 'stdio', 'http', or 'sse'.")
except HTTPException:
@@ -382,7 +386,7 @@ async def discover_tools(tool_id: str):
raise HTTPException(status_code=502, detail=f"Discovery failed: {msg}")
tool_names = [t["name"] for t in raw_tools]
services, service_groups, all_read, all_write = _classify_services(tool_names, tool.name)
services, service_groups, all_read, all_write = classify_services(tool_names, tool.name)
permissions: dict[str, Any] = {n: tool.tool_permissions.get(n, "ask") for n in tool_names}
permissions["_categories"] = {"read": all_read, "write": all_write}
permissions["_services"] = services
@@ -391,7 +395,7 @@ async def discover_tools(tool_id: str):
permissions["_tool_schemas"] = {t["name"]: t.get("inputSchema") for t in raw_tools if t.get("inputSchema")}
tool.tool_permissions = permissions
_save(tool)
p_save(tool)
return {"ok": True, "tool": tool.model_dump()}
@@ -400,7 +404,7 @@ async def discover_tools(tool_id: str):
# Microsoft 365 device-code login (runs in the backend, not the MCP server)
# ---------------------------------------------------------------------------
_m365_login_processes: dict[str, dict] = {} # tool_id -> {proc, device_code, status, email}
P_M365_LOGIN_PROCESSES: dict[str, dict] = {} # tool_id -> {proc, device_code, status, email}
@tools_lib.router.post("/{tool_id}/m365/device-login")
@@ -412,8 +416,8 @@ async def m365_device_login(tool_id: str):
"""
import subprocess
_load(tool_id)
script = _m365_server_script()
p_load(tool_id)
script = m365_server_script()
if not os.path.isfile(script):
raise HTTPException(status_code=500, detail="M365 MCP server not installed")
@@ -427,12 +431,12 @@ async def m365_device_login(tool_id: str):
if not cmd:
raise HTTPException(status_code=500, detail="No node/electron found")
env = {**os.environ, **_m365_cache_env()}
env = {**os.environ, **m365_cache_env()}
if cmd == electron:
env["ELECTRON_RUN_AS_NODE"] = "1"
# Kill any existing login process for this tool
existing = _m365_login_processes.pop(tool_id, None)
existing = P_M365_LOGIN_PROCESSES.pop(tool_id, None)
if existing and existing.get("proc"):
try:
existing["proc"].kill()
@@ -449,7 +453,7 @@ async def m365_device_login(tool_id: str):
import threading
login_state: dict = {"proc": proc, "status": "waiting_for_code", "device_code": "", "device_code_url": "", "email": None, "output": ""}
def _read_output():
def p_read_output():
import re
for line in proc.stdout:
login_state["output"] += line
@@ -469,8 +473,8 @@ async def m365_device_login(tool_id: str):
login_state["status"] = "connected"
# Try to extract email from output
try:
import json as _j
result = _j.loads(login_state["output"].strip().split("\n")[-1])
import json
result = json.loads(login_state["output"].strip().split("\n")[-1])
if result.get("success"):
ud = result.get("userData", {})
login_state["email"] = ud.get("userPrincipalName") or ud.get("displayName")
@@ -478,20 +482,20 @@ async def m365_device_login(tool_id: str):
pass
# Update tool status
try:
t = _load(tool_id)
t = p_load(tool_id)
t.auth_status = "connected"
if login_state.get("email"):
t.connected_account_email = login_state["email"]
_save(t)
p_save(t)
except Exception:
pass
else:
login_state["status"] = "error"
thread = threading.Thread(target=_read_output, daemon=True)
thread = threading.Thread(target=p_read_output, daemon=True)
thread.start()
_m365_login_processes[tool_id] = login_state
P_M365_LOGIN_PROCESSES[tool_id] = login_state
# Wait briefly for device code to appear
for _ in range(30):
@@ -512,13 +516,13 @@ async def m365_device_login(tool_id: str):
@tools_lib.router.get("/{tool_id}/m365/device-login/status")
async def m365_device_login_status(tool_id: str):
"""Poll the status of a pending M365 device-code login."""
state = _m365_login_processes.get(tool_id)
state = P_M365_LOGIN_PROCESSES.get(tool_id)
if not state:
# Check if already connected via cached token
cache_env = _m365_cache_env()
cache_env = m365_cache_env()
cache_path = cache_env["MS365_MCP_TOKEN_CACHE_PATH"]
if os.path.isfile(cache_path):
tool = _load(tool_id)
tool = p_load(tool_id)
if tool.auth_status == "connected":
return {"status": "connected", "email": tool.connected_account_email}
return {"status": "no_login_in_progress"}
@@ -527,10 +531,10 @@ async def m365_device_login_status(tool_id: str):
result: dict = {"status": status}
if status == "connected":
result["email"] = state.get("email")
_m365_login_processes.pop(tool_id, None)
P_M365_LOGIN_PROCESSES.pop(tool_id, None)
elif status == "error":
result["message"] = "Login failed"
_m365_login_processes.pop(tool_id, None)
P_M365_LOGIN_PROCESSES.pop(tool_id, None)
return result
@@ -538,21 +542,21 @@ async def m365_device_login_status(tool_id: str):
@tools_lib.router.post("/{tool_id}/m365/disconnect")
async def m365_disconnect(tool_id: str):
"""Disconnect M365 by clearing the cached token."""
tool = _load(tool_id)
cache_env = _m365_cache_env()
tool = p_load(tool_id)
cache_env = m365_cache_env()
for path in cache_env.values():
if os.path.isfile(path):
os.remove(path)
tool.auth_status = "configured"
tool.connected_account_email = None
_save(tool)
p_save(tool)
return {"ok": True, "tool": tool.model_dump()}
@tools_lib.router.post("/{tool_id}/oauth/disconnect")
async def oauth_disconnect(tool_id: str):
"""Clear OAuth tokens and reset auth status so the user can reconnect with a different account."""
tool = _load(tool_id)
tool = p_load(tool_id)
access_token = tool.oauth_tokens.get("access_token")
if access_token and tool.name.lower() != "notion":
@@ -570,7 +574,7 @@ async def oauth_disconnect(tool_id: str):
tool.oauth_tokens = {}
tool.auth_status = "configured"
tool.connected_account_email = None
_save(tool)
p_save(tool)
return {"ok": True, "tool": tool.model_dump()}
@@ -578,8 +582,8 @@ async def oauth_disconnect(tool_id: str):
async def oauth_start(tool_id: str):
"""Return the OAuth start URL for this tool. All built-in providers
proxy through Fly so client_secret values stay server-side."""
tool = _load(tool_id)
proxied = _proxied_provider_for(tool)
tool = p_load(tool_id)
proxied = proxied_provider_for(tool)
if not proxied:
raise HTTPException(
status_code=400,
@@ -587,11 +591,11 @@ async def oauth_start(tool_id: str):
)
from backend.config.install_id import get_install_id
install_id = get_install_id()
_port = os.environ.get("OPENSWARM_PORT", "8324")
p_port = os.environ.get("OPENSWARM_PORT", "8324")
params = {
"install_id": install_id,
"tool_id": tool_id,
"local_port": _port,
"local_port": p_port,
}
auth_url = (
f"{OPENSWARM_OAUTH_BASE_URL}/api/oauth/{proxied}/start?"
@@ -648,7 +652,7 @@ async def oauth_cloud_claim(
data = resp.json()
tokens = data.get("tokens", {}) or {}
tool = _load(tool_id)
tool = p_load(tool_id)
# Google's token endpoint doesn't include the user's email; fetch it
# from userinfo so the UI can show "you connected you@gmail.com"
# rather than the generic "Google account" placeholder.
@@ -663,9 +667,9 @@ async def oauth_cloud_claim(
tokens["email"] = info_resp.json().get("email") or ""
except Exception as e:
logger.warning("Google userinfo lookup post-claim failed: %s", e)
_persist_cloud_tokens(tool, tokens)
_save(tool)
return _connected_html()
persist_cloud_tokens(tool, tokens)
p_save(tool)
return p_connected_html()
@tools_lib.router.post("/google-oauth-token")
+61 -61
View File
@@ -56,7 +56,7 @@ class FetchBody(BaseModel):
# ---------------------------------------------------------------------------
def _join_text(parts: list[dict[str, Any]]) -> str:
def p_join_text(parts: list[dict[str, Any]]) -> str:
out = []
for p in parts:
if isinstance(p, dict) and p.get("type") == "text":
@@ -69,11 +69,11 @@ def _join_text(parts: list[dict[str, Any]]) -> str:
# ---------------------------------------------------------------------------
GEMINI_API_BASE = "https://generativelanguage.googleapis.com/v1beta"
GEMINI_GROUNDING_MODEL = "gemini-2.5-flash" # cheapest + fastest for grounded calls
P_GEMINI_API_BASE = "https://generativelanguage.googleapis.com/v1beta"
P_GEMINI_GROUNDING_MODEL = "gemini-2.5-flash" # cheapest + fastest for grounded calls
OPENAI_API_BASE = "https://api.openai.com/v1"
OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_preview
P_OPENAI_API_BASE = "https://api.openai.com/v1"
P_OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_preview
# Per-attempt timeouts for the search/fetch cascade. The fast-first ORDERING is
# what fixes the ~75s stall (DDG answers in ~1s so the slow grounded backends are
@@ -82,16 +82,16 @@ OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_pr
# provider (no response at all) gets cut. Grounded native search legitimately
# takes 32-42s (httpx ceiling 45s), so its leash sits at 48s, NOT below 45, or
# we'd clip the slow tail of a valid paid call.
_DDG_ATTEMPT_TIMEOUT = 6.0 # DDG answers <1s; >6s is a network hang, fall through
_GROUNDED_ATTEMPT_TIMEOUT = 48.0 # just above the providers' own 45s httpx timeout
P_DDG_ATTEMPT_TIMEOUT = 6.0 # DDG answers <1s; >6s is a network hang, fall through
P_GROUNDED_ATTEMPT_TIMEOUT = 48.0 # just above the providers' own 45s httpx timeout
# Local httpx + trafilatura fetch of a real page; the fast path for /fetch
# (normal pages return in <2s). Set just above WebFetchTool's own 30s httpx
# ceiling so a valid-but-slow page still completes locally instead of being
# clipped down to a grounded summary; only a truly hung server gets cut.
_LOCAL_FETCH_TIMEOUT = 32.0
P_LOCAL_FETCH_TIMEOUT = 32.0
async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: bool) -> dict:
async def p_gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: bool) -> dict:
"""Call Gemini with googleSearch (+ optionally urlContext) grounding.
Returns {"text": grounded_answer, "chunks": [(title, uri), ...],
@@ -106,7 +106,7 @@ async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: b
"tools": tools,
"generationConfig": {"thinkingConfig": {"thinkingBudget": 0}},
}
url = f"{GEMINI_API_BASE}/models/{GEMINI_GROUNDING_MODEL}:generateContent"
url = f"{P_GEMINI_API_BASE}/models/{P_GEMINI_GROUNDING_MODEL}:generateContent"
async with httpx.AsyncClient(timeout=45.0) as client:
r = await client.post(
url,
@@ -133,7 +133,7 @@ async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: b
return {"text": text, "chunks": chunks, "queries": queries}
def _format_grounded_as_search_results(grounded: dict, query: str) -> str:
def p_format_grounded_as_search_results(grounded: dict, query: str) -> str:
"""Format Gemini grounding output to match WebSearchTool's text shape."""
lines = []
chunks = grounded.get("chunks") or []
@@ -147,7 +147,7 @@ def _format_grounded_as_search_results(grounded: dict, query: str) -> str:
return "\n\n".join(lines)
def _format_grounded_as_fetch(grounded: dict, url: str) -> str:
def p_format_grounded_as_fetch(grounded: dict, url: str) -> str:
"""Format Gemini urlContext output to match WebFetchTool's text shape."""
parts = [f"Contents of {url}:", ""]
text = grounded.get("text") or ""
@@ -161,7 +161,7 @@ def _format_grounded_as_fetch(grounded: dict, url: str) -> str:
return "\n".join(parts)
def _resolve_gemini_api_key() -> str | None:
def p_resolve_gemini_api_key() -> str | None:
"""Pull the AI Studio API key from settings, or None."""
try:
from backend.apps.settings.store import load_settings
@@ -171,7 +171,7 @@ def _resolve_gemini_api_key() -> str | None:
return None
def _resolve_openai_api_key() -> str | None:
def p_resolve_openai_api_key() -> str | None:
try:
from backend.apps.settings.store import load_settings
s = load_settings()
@@ -184,38 +184,38 @@ def _resolve_openai_api_key() -> str | None:
# `_refresh_9r_connected()` rather than hit on every search call ,
# 9Router's /api/providers is fast but not free, and we already
# query it from many places.
_NINE_ROUTER_CONNECTED: set[str] = set()
_NINE_ROUTER_CACHE_AT: float = 0.0
P_NINE_ROUTER_CONNECTED: set[str] = set()
P_NINE_ROUTER_CACHE_AT: float = 0.0
async def _refresh_9r_connected() -> set[str]:
async def p_refresh_9r_connected() -> set[str]:
"""Return the set of currently-active 9Router subscription providers
(e.g. {"claude", "codex", "antigravity", "gemini-cli"}). Cached for
20s to keep search/fetch endpoints snappy."""
global _NINE_ROUTER_CONNECTED, _NINE_ROUTER_CACHE_AT
global P_NINE_ROUTER_CONNECTED, P_NINE_ROUTER_CACHE_AT
import time
now = time.time()
if now - _NINE_ROUTER_CACHE_AT < 20.0:
return _NINE_ROUTER_CONNECTED
if now - P_NINE_ROUTER_CACHE_AT < 20.0:
return P_NINE_ROUTER_CONNECTED
try:
from backend.apps.nine_router.process import is_running, get_providers
if not is_running():
_NINE_ROUTER_CONNECTED = set()
P_NINE_ROUTER_CONNECTED = set()
else:
conns = await get_providers()
_NINE_ROUTER_CONNECTED = {
P_NINE_ROUTER_CONNECTED = {
c.get("provider")
for c in conns
if isinstance(c, dict) and c.get("isActive") and c.get("provider")
}
_NINE_ROUTER_CACHE_AT = now
P_NINE_ROUTER_CACHE_AT = now
except Exception:
# Cache stays; best-effort.
pass
return _NINE_ROUTER_CONNECTED
return P_NINE_ROUTER_CONNECTED
async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> dict:
async def p_gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> dict:
"""Call 9Router's /v1/messages endpoint with a Gemini model so the
user's OAuth subscription (Gemini CLI or Antigravity) covers the
search call instead of needing a separate AI Studio API key.
@@ -223,12 +223,12 @@ async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> di
Routes through Anthropic-shape against 9Router's translator. We
request a tool result naturally; the translator surfaces grounded
URIs as text + cited sources in the response body. Format-shape
matches the existing `_gemini_grounded_call` so downstream
matches the existing `p_gemini_grounded_call` so downstream
`_format_grounded_as_search_results` works unchanged."""
import httpx
# Prefer Gemini CLI (broader model coverage). Fall back to
# Antigravity if CLI isn't connected.
connected = await _refresh_9r_connected()
connected = await p_refresh_9r_connected()
if "gemini-cli" in connected:
model = "gc/gemini-2.5-flash"
elif "antigravity" in connected:
@@ -269,12 +269,12 @@ async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> di
return {"text": text, "chunks": []}
async def _openai_websearch_via_9router(query: str) -> dict:
async def p_openai_websearch_via_9router(query: str) -> dict:
"""Same idea, but for OpenAI's web_search_preview through Codex's
9Router connection. Goes through 9Router's openai-compat endpoint
(the responses API) so the user's Codex subscription covers it."""
import httpx
connected = await _refresh_9r_connected()
connected = await p_refresh_9r_connected()
if "codex" not in connected:
return {}
body = {
@@ -302,20 +302,20 @@ async def _openai_websearch_via_9router(query: str) -> dict:
return {"text": text, "chunks": []}
async def _openai_websearch(api_key: str, query: str) -> dict:
async def p_openai_websearch(api_key: str, query: str) -> dict:
"""Call OpenAI Responses API with the web_search_preview tool.
Returns {"text": grounded_answer, "chunks": [(title, uri), ...]}.
"""
import httpx
body = {
"model": OPENAI_SEARCH_MODEL,
"model": P_OPENAI_SEARCH_MODEL,
"input": f"Search the web for: {query}\n\nReturn a concise summary. Cite sources.",
"tools": [{"type": "web_search_preview"}],
}
async with httpx.AsyncClient(timeout=45.0) as client:
r = await client.post(
f"{OPENAI_API_BASE}/responses",
f"{P_OPENAI_API_BASE}/responses",
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=body,
)
@@ -341,20 +341,20 @@ async def _openai_websearch(api_key: str, query: str) -> dict:
return {"text": "".join(text_parts), "chunks": chunks, "queries": [query]}
async def _openai_urlfetch(api_key: str, url: str, prompt: str | None) -> dict:
async def p_openai_urlfetch(api_key: str, url: str, prompt: str | None) -> dict:
"""Use OpenAI's web_search_preview to fetch/summarize a specific URL."""
prompt_text = f"Fetch and summarize the content at: {url}"
if prompt:
prompt_text += f"\n\nFocus on: {prompt}"
import httpx
body = {
"model": OPENAI_SEARCH_MODEL,
"model": P_OPENAI_SEARCH_MODEL,
"input": prompt_text,
"tools": [{"type": "web_search_preview"}],
}
async with httpx.AsyncClient(timeout=45.0) as client:
r = await client.post(
f"{OPENAI_API_BASE}/responses",
f"{P_OPENAI_API_BASE}/responses",
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=body,
)
@@ -391,8 +391,8 @@ async def search(body: SearchBody) -> dict:
If the primary's own native path fails, we cascade to whichever
other provider's credentials are available, then DDG last."""
gemini_key = _resolve_gemini_api_key()
openai_key = _resolve_openai_api_key()
gemini_key = p_resolve_gemini_api_key()
openai_key = p_resolve_openai_api_key()
primary = (body.primary or "").lower()
errors: list[str] = []
@@ -403,20 +403,20 @@ async def search(body: SearchBody) -> dict:
f"Search the web for: {body.query}\n\n"
f"Return a concise summary of what you found. Cite sources."
)
grounded = await _gemini_grounded_call(gemini_key, prompt, use_url_context=False)
grounded = await p_gemini_grounded_call(gemini_key, prompt, use_url_context=False)
return {
"query": body.query,
"results": _format_grounded_as_search_results(grounded, body.query),
"results": p_format_grounded_as_search_results(grounded, body.query),
"backend": "gemini_native",
}
async def try_openai():
if not openai_key:
return None
grounded = await _openai_websearch(openai_key, body.query)
grounded = await p_openai_websearch(openai_key, body.query)
return {
"query": body.query,
"results": _format_grounded_as_search_results(grounded, body.query),
"results": p_format_grounded_as_search_results(grounded, body.query),
"backend": "openai_native",
}
@@ -425,22 +425,22 @@ async def search(body: SearchBody) -> dict:
f"Search the web for: {body.query}\n\n"
f"Return a concise summary of what you found. Cite sources."
)
grounded = await _gemini_grounded_via_9router(prompt, use_url_context=False)
grounded = await p_gemini_grounded_via_9router(prompt, use_url_context=False)
if not grounded.get("text"):
return None
return {
"query": body.query,
"results": _format_grounded_as_search_results(grounded, body.query),
"results": p_format_grounded_as_search_results(grounded, body.query),
"backend": "gemini_subscription",
}
async def try_openai_subscription():
grounded = await _openai_websearch_via_9router(body.query)
grounded = await p_openai_websearch_via_9router(body.query)
if not grounded.get("text"):
return None
return {
"query": body.query,
"results": _format_grounded_as_search_results(grounded, body.query),
"results": p_format_grounded_as_search_results(grounded, body.query),
"backend": "openai_subscription",
}
@@ -473,8 +473,8 @@ async def search(body: SearchBody) -> dict:
if primary == "openai":
grounded = grounded[2:] + grounded[:2]
cascade = [("ddg", try_ddg, _DDG_ATTEMPT_TIMEOUT)] + [
(name, fn, _GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
cascade = [("ddg", try_ddg, P_DDG_ATTEMPT_TIMEOUT)] + [
(name, fn, P_GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
]
for name, fn, timeout in cascade:
@@ -490,7 +490,7 @@ async def search(body: SearchBody) -> dict:
errors.append(f"{name}: {str(e)[:150]}")
# Everything failed. Be honest about why instead of an empty "no results".
connected = await _refresh_9r_connected()
connected = await p_refresh_9r_connected()
has_subscription = bool(connected & {"codex", "antigravity", "gemini-cli"})
if not (gemini_key or openai_key or has_subscription):
tail = (
@@ -524,8 +524,8 @@ async def fetch(body: FetchBody) -> dict:
except SSRFBlocked as exc:
from fastapi import HTTPException
raise HTTPException(status_code=400, detail=f"Refused: {exc}")
gemini_key = _resolve_gemini_api_key()
openai_key = _resolve_openai_api_key()
gemini_key = p_resolve_gemini_api_key()
openai_key = p_resolve_openai_api_key()
primary = (body.primary or "").lower()
async def try_gemini():
@@ -534,22 +534,22 @@ async def fetch(body: FetchBody) -> dict:
prompt_bits = [f"Fetch and summarize this URL: {body.url}"]
if body.prompt:
prompt_bits.append(f"Focus on: {body.prompt}")
grounded = await _gemini_grounded_call(
grounded = await p_gemini_grounded_call(
gemini_key, "\n".join(prompt_bits), use_url_context=True,
)
return {
"url": body.url,
"content": _format_grounded_as_fetch(grounded, body.url),
"content": p_format_grounded_as_fetch(grounded, body.url),
"backend": "gemini_native",
}
async def try_openai():
if not openai_key:
return None
grounded = await _openai_urlfetch(openai_key, body.url, body.prompt)
grounded = await p_openai_urlfetch(openai_key, body.url, body.prompt)
return {
"url": body.url,
"content": _format_grounded_as_fetch(grounded, body.url),
"content": p_format_grounded_as_fetch(grounded, body.url),
"backend": "openai_native",
}
@@ -557,14 +557,14 @@ async def fetch(body: FetchBody) -> dict:
prompt_bits = [f"Fetch and summarize this URL: {body.url}"]
if body.prompt:
prompt_bits.append(f"Focus on: {body.prompt}")
grounded = await _gemini_grounded_via_9router(
grounded = await p_gemini_grounded_via_9router(
"\n".join(prompt_bits), use_url_context=True,
)
if not grounded.get("text"):
return None
return {
"url": body.url,
"content": _format_grounded_as_fetch(grounded, body.url),
"content": p_format_grounded_as_fetch(grounded, body.url),
"backend": "gemini_subscription",
}
@@ -574,12 +574,12 @@ async def fetch(body: FetchBody) -> dict:
prompt = f"Fetch this URL and summarize: {body.url}"
if body.prompt:
prompt += f"\nFocus on: {body.prompt}"
grounded = await _openai_websearch_via_9router(prompt)
grounded = await p_openai_websearch_via_9router(prompt)
if not grounded.get("text"):
return None
return {
"url": body.url,
"content": _format_grounded_as_fetch(grounded, body.url),
"content": p_format_grounded_as_fetch(grounded, body.url),
"backend": "openai_subscription",
}
@@ -597,7 +597,7 @@ async def fetch(body: FetchBody) -> dict:
parts = await WebFetchTool().execute(
{"url": body.url, "prompt": body.prompt or ""}, None,
)
text = _join_text(parts)
text = p_join_text(parts)
local_text = text
if text.startswith(("HTTP error", "Error fetching", "Refused to fetch")):
return None
@@ -615,8 +615,8 @@ async def fetch(body: FetchBody) -> dict:
if primary == "openai":
grounded = grounded[2:] + grounded[:2]
cascade = [("local", try_local, _LOCAL_FETCH_TIMEOUT)] + [
(name, fn, _GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
cascade = [("local", try_local, P_LOCAL_FETCH_TIMEOUT)] + [
(name, fn, P_GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
]
errors: list[str] = []
+55 -49
View File
@@ -10,10 +10,10 @@ from backend.config.paths import AUTH_TOKEN_FILE
logger = logging.getLogger(__name__)
_TOKEN: str = ""
TOKEN: str = ""
def _write_atomic(path: str, data: str, mode: int = 0o600) -> None:
def p_write_atomic(path: str, data: str, mode: int = 0o600) -> None:
"""Atomic write to `path` at the given file mode; never world-readable or half-written."""
os.makedirs(os.path.dirname(path), exist_ok=True)
tmp = path + ".tmp"
@@ -29,100 +29,102 @@ def _write_atomic(path: str, data: str, mode: int = 0o600) -> None:
os.replace(tmp, path)
# Public - called by main.py
def init_auth_token() -> str:
"""Load the per-install token from disk, or mint one if missing; reused across restarts so Electron's cached copy stays valid."""
global _TOKEN
global TOKEN
try:
if os.path.exists(AUTH_TOKEN_FILE):
with open(AUTH_TOKEN_FILE, "r", encoding="utf-8") as f:
existing = f.read().strip()
if existing and 16 <= len(existing) <= 512:
_TOKEN = existing
TOKEN = existing
logger.info(
f"auth: reusing existing token from {AUTH_TOKEN_FILE}"
)
return _TOKEN
return TOKEN
except Exception as e:
logger.warning(f"auth: failed to read existing token, generating new: {e}")
_TOKEN = secrets.token_urlsafe(32)
TOKEN = secrets.token_urlsafe(32)
try:
_write_atomic(AUTH_TOKEN_FILE, _TOKEN, mode=0o600)
p_write_atomic(AUTH_TOKEN_FILE, TOKEN, mode=0o600)
logger.info(f"auth: wrote token to {AUTH_TOKEN_FILE} (mode 0600)")
except Exception as e:
# If we can't write the file, Electron can't read it; log loudly but don't crash.
logger.error(f"auth: failed to write token file: {e}")
return _TOKEN
return TOKEN
# Public - called by main.py
def get_auth_token() -> str:
"""Return the current token. Empty string if init_auth_token() hasn't run."""
return _TOKEN
return TOKEN
class _TokenScrubFilter(logging.Filter):
class P_TokenScrubFilter(logging.Filter):
"""Logging filter that redacts the install token from log records (defense in depth)."""
_PLACEHOLDER = "<REDACTED:openswarm-token>"
P_PLACEHOLDER = "<REDACTED:openswarm-token>"
@staticmethod
def _args_might_contain_token(args) -> bool:
def p_args_might_contain_token(args) -> bool:
"""Cheap pre-check; avoids eager %-formatting on the >99% of records that don't mention the token."""
if not args:
return False
items = args if isinstance(args, (tuple, list)) else (args,)
for a in items:
if isinstance(a, str) and _TOKEN in a:
if isinstance(a, str) and TOKEN in a:
return True
if isinstance(a, dict):
for v in a.values():
if isinstance(v, str) and _TOKEN in v:
if isinstance(v, str) and TOKEN in v:
return True
return False
@classmethod
def _scrub_args(cls, args):
def p_scrub_args(cls, args):
"""Scrub token from args while preserving tuple/dict shape; uvicorn's AccessFormatter unpacks args as a 5-tuple and explodes on None."""
if args is None:
return args
if isinstance(args, dict):
new_dict = None
for k, v in args.items():
if isinstance(v, str) and _TOKEN in v:
if isinstance(v, str) and TOKEN in v:
if new_dict is None:
new_dict = dict(args)
new_dict[k] = v.replace(_TOKEN, cls._PLACEHOLDER)
new_dict[k] = v.replace(TOKEN, cls.P_PLACEHOLDER)
return new_dict if new_dict is not None else args
if isinstance(args, tuple):
new_list = None
for i, v in enumerate(args):
if isinstance(v, str) and _TOKEN in v:
if isinstance(v, str) and TOKEN in v:
if new_list is None:
new_list = list(args)
new_list[i] = v.replace(_TOKEN, cls._PLACEHOLDER)
new_list[i] = v.replace(TOKEN, cls.P_PLACEHOLDER)
return tuple(new_list) if new_list is not None else args
if isinstance(args, str) and _TOKEN in args:
return args.replace(_TOKEN, cls._PLACEHOLDER)
if isinstance(args, str) and TOKEN in args:
return args.replace(TOKEN, cls.P_PLACEHOLDER)
return args
def filter(self, record: logging.LogRecord) -> bool: # pragma: no cover (defensive)
if not _TOKEN:
if not TOKEN:
return True
# Fast path: skip eager %-formatting on records that don't mention the token.
raw_msg = record.msg if isinstance(record.msg, str) else ""
if _TOKEN not in raw_msg and not self._args_might_contain_token(record.args):
if TOKEN not in raw_msg and not self.p_args_might_contain_token(record.args):
return True
# Slow path: in-place args rewrite (preserves shape for AccessFormatter), then re-render to catch tokens buried in custom reprs.
try:
if isinstance(record.msg, str) and _TOKEN in record.msg:
record.msg = record.msg.replace(_TOKEN, self._PLACEHOLDER)
if isinstance(record.msg, str) and TOKEN in record.msg:
record.msg = record.msg.replace(TOKEN, self.P_PLACEHOLDER)
scrubbed = self._scrub_args(record.args)
if scrubbed is not record.args:
record.args = scrubbed
try:
rendered = record.getMessage()
if _TOKEN in rendered:
record.msg = rendered.replace(_TOKEN, self._PLACEHOLDER)
if TOKEN in rendered:
record.msg = rendered.replace(TOKEN, self.P_PLACEHOLDER)
record.args = None
except Exception:
pass
@@ -132,19 +134,20 @@ class _TokenScrubFilter(logging.Filter):
return True
_scrubber_installed = False
P_SCRUBBER_INSTALLED = False
# Public - called by main.py
def install_token_scrubber() -> None:
"""Attach the scrubbing filter to every existing AND future log handler; logger-level filters miss propagated child records."""
global _scrubber_installed
if _scrubber_installed:
global P_SCRUBBER_INSTALLED
if P_SCRUBBER_INSTALLED:
return
scrubber = _TokenScrubFilter()
scrubber = P_TokenScrubFilter()
def _attach(handler: logging.Handler) -> None:
if not any(isinstance(f, _TokenScrubFilter) for f in handler.filters):
def attach(handler: logging.Handler) -> None:
if not any(isinstance(f, P_TokenScrubFilter) for f in handler.filters):
handler.addFilter(scrubber)
loggers: list[logging.Logger] = [logging.getLogger()]
@@ -153,26 +156,26 @@ def install_token_scrubber() -> None:
loggers.append(logger)
for logger in loggers:
for h in list(logger.handlers):
_attach(h)
attach(h)
# Patch addHandler so handlers attached later (uvicorn finishes log config after main.py imports) get the scrubber too.
_original_addHandler = logging.Logger.addHandler
original_addHandler = logging.Logger.addHandler
def _patched_addHandler(self: logging.Logger, hdlr: logging.Handler) -> None:
_attach(hdlr)
return _original_addHandler(self, hdlr)
def patched_addHandler(self: logging.Logger, hdlr: logging.Handler) -> None:
attach(hdlr)
return original_addHandler(self, hdlr)
logging.Logger.addHandler = _patched_addHandler # type: ignore[assignment]
logging.Logger.addHandler = patched_addHandler # type: ignore[assignment]
root = logging.getLogger()
if not any(isinstance(f, _TokenScrubFilter) for f in root.filters):
if not any(isinstance(f, P_TokenScrubFilter) for f in root.filters):
root.addFilter(scrubber)
_scrubber_installed = True
P_SCRUBBER_INSTALLED = True
# Auth-exempt paths: external redirects with their own nonce/state validation, plus the bootstrap health probe.
_AUTH_EXEMPT_EXACT = {
P_AUTH_EXEMPT_EXACT = {
"/api/subscriptions/callback",
"/api/tools/oauth/callback",
"/api/tools/oauth/cloud-claim",
@@ -191,7 +194,7 @@ _AUTH_EXEMPT_EXACT = {
"/api/dev/token",
}
_AUTH_EXEMPT_PREFIX = (
P_AUTH_EXEMPT_PREFIX = (
# Electron polls /api/health/check before loading the token.
"/api/health",
# 9Router proxies OpenAI requests with the user's sk-... bearer, not our local token; localhost-only is the gate.
@@ -203,11 +206,12 @@ _AUTH_EXEMPT_PREFIX = (
)
# Public - called by main.py
def is_path_exempt(path: str) -> bool:
"""True if this request path bypasses token auth."""
if path in _AUTH_EXEMPT_EXACT:
if path in P_AUTH_EXEMPT_EXACT:
return True
for p in _AUTH_EXEMPT_PREFIX:
for p in P_AUTH_EXEMPT_PREFIX:
if path.startswith(p):
return True
return False
@@ -224,9 +228,10 @@ def extract_bearer(header_value: str | None) -> str:
return ""
# Public - called by main.py
def request_matches_token(request_headers: dict, query_params: dict | None = None) -> bool:
"""Validate that an HTTP/WS request carries our token (Bearer, x-openswarm-token, or ?token=); constant-time compare."""
if not _TOKEN:
if not TOKEN:
# Backend not initialized: fail closed. Only test fixtures that bypass main hit this.
return False
@@ -250,13 +255,13 @@ def request_matches_token(request_headers: dict, query_params: dict | None = Non
candidates.append(qp_token)
for candidate in candidates:
if secrets.compare_digest(candidate, _TOKEN):
if secrets.compare_digest(candidate, TOKEN):
return True
return False
# WS Origin allowlist: Electron packaged is file://, dev is localhost:3000, some Electron contexts send bare "null".
_ORIGIN_ALLOWLIST_DEV = {
P_ORIGIN_ALLOWLIST_DEV = {
"http://localhost:3000",
"http://127.0.0.1:3000",
"file://",
@@ -264,12 +269,13 @@ _ORIGIN_ALLOWLIST_DEV = {
}
# Public - called by main.py
def is_origin_allowed(origin: str | None) -> bool:
"""True if the WS connection's Origin header is from our app."""
if origin is None:
# Native WS client / curl / MCP subprocess: token check still required, so allow.
return True
if origin in _ORIGIN_ALLOWLIST_DEV:
if origin in P_ORIGIN_ALLOWLIST_DEV:
return True
# Packaged Electron file:// includes paths like file:///Applications/OpenSwarm.app/...; match by prefix.
if origin.startswith("file://"):
+2 -2
View File
@@ -32,8 +32,8 @@ class MainApp:
for sub_app in sub_apps:
debug(sub_app.name)
await stack.enter_async_context(sub_app.lifespan())
_port = os.environ.get("OPENSWARM_PORT", "8324")
print(f"\nCheck out the API docs at: http://127.0.0.1:{_port}/docs\n")
port = os.environ.get("OPENSWARM_PORT", "8324")
print(f"\nCheck out the API docs at: http://127.0.0.1:{port}/docs\n")
yield
self.app = FastAPI(lifespan=lifespan)
+15 -14
View File
@@ -7,39 +7,40 @@ import uuid
from backend.config.paths import DATA_ROOT
_INSTALL_ID_FILE = os.path.join(DATA_ROOT, "install_id")
_cached: str | None = None
P_INSTALL_ID_FILE = os.path.join(DATA_ROOT, "install_id")
P_CACHED: str | None = None
# Public - called by tools_lib.py, mcp_config.py
def get_install_id() -> str:
"""Return the persistent install_id, generating and persisting on first call."""
global _cached
if _cached:
return _cached
global P_CACHED
if P_CACHED:
return P_CACHED
try:
with open(_INSTALL_ID_FILE, "r", encoding="utf-8") as f:
with open(P_INSTALL_ID_FILE, "r", encoding="utf-8") as f:
existing = f.read().strip()
if _looks_like_uuid(existing):
_cached = existing
return _cached
if p_looks_like_uuid(existing):
P_CACHED = existing
return P_CACHED
except FileNotFoundError:
pass
except Exception:
pass
fresh = str(uuid.uuid4())
os.makedirs(os.path.dirname(_INSTALL_ID_FILE) or ".", exist_ok=True)
fd = os.open(_INSTALL_ID_FILE, os.O_CREAT | os.O_WRONLY | os.O_TRUNC, 0o600)
os.makedirs(os.path.dirname(P_INSTALL_ID_FILE) or ".", exist_ok=True)
fd = os.open(P_INSTALL_ID_FILE, os.O_CREAT | os.O_WRONLY | os.O_TRUNC, 0o600)
try:
os.write(fd, fresh.encode("utf-8"))
finally:
os.close(fd)
_cached = fresh
return _cached
P_CACHED = fresh
return P_CACHED
def _looks_like_uuid(s: str) -> bool:
def p_looks_like_uuid(s: str) -> bool:
if len(s) != 36:
return False
try:
+2 -2
View File
@@ -21,12 +21,12 @@ import time
logger = logging.getLogger(__name__)
_write_lock = threading.Lock()
P_WRITE_LOCK = threading.Lock()
def atomic_write_json(path: str, payload, *, indent: int = 2) -> None:
directory = os.path.dirname(path) or "."
with _write_lock:
with P_WRITE_LOCK:
os.makedirs(directory, exist_ok=True)
fd, tmp = tempfile.mkstemp(prefix=".tmp-", suffix=".json", dir=directory)
try:
+9 -9
View File
@@ -3,20 +3,20 @@
import os
import sys
_BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
P_BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
_is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
P_IS_PACKAGED = os.environ.get("OPENSWARM_PACKAGED") == "1"
if _is_packaged:
if P_IS_PACKAGED:
if sys.platform == "darwin":
_app_support = os.path.join(os.path.expanduser("~"), "Library", "Application Support", "OpenSwarm")
app_support = os.path.join(os.path.expanduser("~"), "Library", "Application Support", "OpenSwarm")
elif sys.platform == "win32":
_app_support = os.path.join(os.environ.get("APPDATA", os.path.expanduser("~")), "OpenSwarm")
app_support = os.path.join(os.environ.get("APPDATA", os.path.expanduser("~")), "OpenSwarm")
else:
_app_support = os.path.join(os.environ.get("XDG_DATA_HOME", os.path.join(os.path.expanduser("~"), ".local", "share")), "OpenSwarm")
DATA_ROOT = os.path.join(_app_support, "data")
app_support = os.path.join(os.environ.get("XDG_DATA_HOME", os.path.join(os.path.expanduser("~"), ".local", "share")), "OpenSwarm")
DATA_ROOT = os.path.join(app_support, "data")
else:
DATA_ROOT = os.path.join(_BACKEND_DIR, "data")
DATA_ROOT = os.path.join(P_BACKEND_DIR, "data")
SESSIONS_DIR = os.path.join(DATA_ROOT, "sessions")
TOOLS_DIR = os.path.join(DATA_ROOT, "tools")
@@ -33,4 +33,4 @@ TRUSTED_SENSITIVE_PATHS_PATH = os.path.join(DATA_ROOT, "trusted_sensitive_paths.
# Per-install auth token for the localhost API; see auth.py.
AUTH_TOKEN_FILE = os.path.join(DATA_ROOT, "auth.token")
BACKEND_DIR = _BACKEND_DIR
BACKEND_DIR = P_BACKEND_DIR
+50 -51
View File
@@ -8,13 +8,13 @@ from uuid import uuid4
# invisible because nothing configured the 'backend' logger; every debugging
# session re-paid that blindness. Idempotent so uvicorn reloads don't stack
# handlers; uvicorn's own access logs are untouched.
_backend_logger = logging.getLogger("backend")
if not _backend_logger.handlers:
_h = logging.StreamHandler()
_h.setFormatter(logging.Formatter("%(asctime)s %(levelname).1s %(name)s: %(message)s", "%H:%M:%S"))
_backend_logger.addHandler(_h)
_backend_logger.setLevel(logging.INFO)
_backend_logger.propagate = False
P_BACKEND_LOGGER = logging.getLogger("backend")
if not P_BACKEND_LOGGER.handlers:
h = logging.StreamHandler()
h.setFormatter(logging.Formatter("%(asctime)s %(levelname).1s %(name)s: %(message)s", "%H:%M:%S"))
P_BACKEND_LOGGER.addHandler(h)
P_BACKEND_LOGGER.setLevel(logging.INFO)
P_BACKEND_LOGGER.propagate = False
logger = logging.getLogger(__name__)
@@ -74,12 +74,12 @@ install_token_scrubber()
# carries it. Platform-agnostic; wrapped so a settings hiccup never blocks
# startup, and the lazy path stays as a fallback.
try:
import uuid as _uuid
from backend.apps.settings.store import load_settings as _load_boot_settings, save_settings as _save_boot_settings
_boot_settings = _load_boot_settings()
if not getattr(_boot_settings, "installation_id", None):
_boot_settings.installation_id = _uuid.uuid4().hex
_save_boot_settings(_boot_settings)
import uuid
from backend.apps.settings.store import load_settings, save_settings
boot_settings = load_settings()
if not getattr(boot_settings, "installation_id", None):
boot_settings.installation_id = uuid.uuid4().hex
save_settings(boot_settings)
except Exception:
pass
@@ -114,7 +114,7 @@ app.add_middleware(
@app.middleware("http")
async def _auth_middleware(request: Request, call_next):
async def auth_middleware(request: Request, call_next):
"""Reject HTTP requests without our per-install bearer token.
Exemptions (see `auth.is_path_exempt`):
@@ -147,9 +147,9 @@ async def _auth_middleware(request: Request, call_next):
# /api/outputs/.../serve/index.html via <iframe src="...">.
auth_ok = request_matches_token(headers, query_params=dict(request.query_params))
if not auth_ok and x_api_key:
import secrets as _s
from backend.auth import get_auth_token as _gt
auth_ok = _s.compare_digest(x_api_key, _gt() or "\x00")
import secrets
from backend.auth import get_auth_token
auth_ok = secrets.compare_digest(x_api_key, get_auth_token() or "\x00")
if not auth_ok:
logger.warning(
f"auth: rejecting {request.method} {request.url.path} "
@@ -184,7 +184,7 @@ async def websocket_session(websocket: WebSocket, session_id: str):
things that end a run are: natural completion, explicit
`agent:stop`, REST `/close`, or process shutdown.
"""
if not _ws_auth_ok(websocket):
if not p_ws_auth_ok(websocket):
return
await ws_manager.connect_session(session_id, websocket)
try:
@@ -257,7 +257,7 @@ async def websocket_session(websocket: WebSocket, session_id: str):
# the agent task, that's intentional. See module docstring.
ws_manager.disconnect_session(session_id, websocket)
def _ws_auth_ok(websocket: WebSocket) -> bool:
def p_ws_auth_ok(websocket: WebSocket) -> bool:
"""Validate token + origin before accepting a WS. Returns True if OK.
On failure closes with 4401 (custom app-level code) and returns False,
@@ -273,8 +273,7 @@ def _ws_auth_ok(websocket: WebSocket) -> bool:
logger.warning(f"ws: rejecting connection to {websocket.url.path}, {reason}")
# Can't `await websocket.close()` before accept(), so schedule the
# close in a task. The client receives a 403 on handshake.
import asyncio as _asyncio
_asyncio.create_task(websocket.close(code=4401))
asyncio.create_task(websocket.close(code=4401))
return False
return True
@@ -285,7 +284,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
pane. On connect we replay the runtime's ring buffer so a Terminal
tab opened mid-session sees the context it missed, then we tail
every subsequent line until disconnect."""
if not _ws_auth_ok(websocket):
if not p_ws_auth_ok(websocket):
return
await websocket.accept()
from backend.apps.outputs.runtime import RUNTIME_MANAGER
@@ -318,15 +317,15 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
# primed with existing lines before we enter the loop.
queue: asyncio.Queue[tuple[str, str]] = asyncio.Queue()
def _on_line(line) -> None:
def on_line(line) -> None:
try:
queue.put_nowait((line.stream, line.text))
except asyncio.QueueFull:
pass
unsubscribe = rt.subscribe(_on_line)
unsubscribe = rt.subscribe(on_line)
def _build_status_frame() -> dict:
def build_status_frame() -> dict:
return {
"event": "runtime:status",
"workspace_id": workspace_id,
@@ -346,7 +345,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
# new-mode preview pointer (Vite dev server); `backend_url` is
# the workspace's optional FastAPI backend (old-mode backend.py
# OR new-mode post-backend_init.sh).
await websocket.send_text(json.dumps(_build_status_frame()))
await websocket.send_text(json.dumps(build_status_frame()))
while True:
stream, text = await queue.get()
await websocket.send_text(json.dumps({
@@ -361,7 +360,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
# to the Vite URL and the preview pane has to know to
# switch over. Re-push status after every runtime line.
if stream == "runtime":
await websocket.send_text(json.dumps(_build_status_frame()))
await websocket.send_text(json.dumps(build_status_frame()))
except WebSocketDisconnect:
pass
finally:
@@ -370,7 +369,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
@app.websocket("/ws/dashboard")
async def websocket_dashboard(websocket: WebSocket):
if not _ws_auth_ok(websocket):
if not p_ws_auth_ok(websocket):
return
await ws_manager.connect_global(websocket)
try:
@@ -446,7 +445,7 @@ async def subscriptions_pending(state: str):
}, headers={"Access-Control-Allow-Origin": "*"})
_SUCCESS_HTML = (
P_SUCCESS_HTML = (
'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif">'
'<div style="text-align:center">'
'<div style="width:64px;height:64px;border-radius:50%;background:#22c55e20;display:flex;align-items:center;justify-content:center;margin:0 auto 16px;font-size:32px">&#10003;</div>'
@@ -490,7 +489,7 @@ async def subscriptions_callback(request: Request):
# Chrome's prefetcher and some extensions speculatively GET URLs.
if state and state in COMPLETED_OAUTH:
logger.info(f"Duplicate OAuth callback for state {state[:8]}... (already completed)")
return HTMLResponse(_SUCCESS_HTML)
return HTMLResponse(P_SUCCESS_HTML)
logger.warning(f"OAuth callback with unknown state {state[:8] if state else '(empty)'}...")
return HTMLResponse('<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Session expired</h2><p style="color:#888">Please try connecting again.</p></div></body></html>')
@@ -508,7 +507,7 @@ async def subscriptions_callback(request: Request):
mark_oauth_completed(state)
logger.info(f"OAuth exchange succeeded for provider={pending.get('provider')}")
return HTMLResponse(_SUCCESS_HTML)
return HTMLResponse(P_SUCCESS_HTML)
@app.post("/api/browser-agent/run")
@@ -550,7 +549,7 @@ async def mcp_meta(action: str, request: Request):
valid options instead of activating (anti-hallucination).
"""
from backend.apps.agents.agent_manager import agent_manager
from backend.apps.tools_lib.tools_lib import _load_all as load_all_tools
from backend.apps.tools_lib.tools_lib import load_all_tools
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
body = await request.json()
@@ -562,7 +561,7 @@ async def mcp_meta(action: str, request: Request):
# sanitized server names; values are extra search-hint tokens appended
# to the haystack. Only generic synonyms, anything that's already in
# the description doesn't need to be listed.
_SERVER_SEARCH_ALIASES: dict[str, list[str]] = {
P_SERVER_SEARCH_ALIASES: dict[str, list[str]] = {
"google-workspace": [
"email", "inbox", "mail", "gmail", "calendar", "schedule",
"events", "drive", "docs", "sheets", "spreadsheet", "slides",
@@ -582,7 +581,7 @@ async def mcp_meta(action: str, request: Request):
"youtube": ["video", "transcript", "channel"],
}
def _connected_servers() -> list[dict]:
def connected_servers() -> list[dict]:
out = []
for t in load_all_tools():
if not (t.mcp_config and t.enabled and t.auth_status in ("configured", "connected")):
@@ -597,7 +596,7 @@ async def mcp_meta(action: str, request: Request):
action_names = [str(k) for k in td.keys() if not str(k).startswith("_")]
except Exception:
pass
aliases = _SERVER_SEARCH_ALIASES.get(sanitized, [])
aliases = P_SERVER_SEARCH_ALIASES.get(sanitized, [])
out.append({
"name": sanitized,
"description": (t.description or "").strip() or f"{t.name} integration",
@@ -606,20 +605,20 @@ async def mcp_meta(action: str, request: Request):
})
return out
def _strip_extras(s: dict) -> dict:
def strip_extras(s: dict) -> dict:
return {k: v for k, v in s.items() if not k.startswith("_")}
if action == "list":
servers = _connected_servers()
servers = connected_servers()
session = agent_manager.sessions.get(parent_session_id) if parent_session_id else None
active_set = set(session.active_mcps) if session else set()
active = [{**_strip_extras(s), "status": "active"} for s in servers if s["name"] in active_set]
available = [{**_strip_extras(s), "status": "available"} for s in servers if s["name"] not in active_set]
active = [{**strip_extras(s), "status": "active"} for s in servers if s["name"] in active_set]
available = [{**strip_extras(s), "status": "available"} for s in servers if s["name"] not in active_set]
return JSONResponse({"active": active, "available": available})
if action == "search":
query = (body.get("query") or "").strip().lower()
servers = _connected_servers()
servers = connected_servers()
session = agent_manager.sessions.get(parent_session_id) if parent_session_id else None
active_set = set(session.active_mcps) if session else set()
# Ranking: substring hits across name+description+sub-tool names+
@@ -644,7 +643,7 @@ async def mcp_meta(action: str, request: Request):
else:
score += 1
if score:
annotated = {**_strip_extras(s), "status": "active" if s["name"] in active_set else "available"}
annotated = {**strip_extras(s), "status": "active" if s["name"] in active_set else "available"}
scored.append((score, annotated))
scored.sort(key=lambda t: (-t[0], 0 if t[1]["status"] == "active" else 1, t[1]["name"]))
matches = [s for _, s in scored[:5]]
@@ -660,7 +659,7 @@ async def mcp_meta(action: str, request: Request):
if not session:
return JSONResponse({"error": "session not found"}, status_code=404)
servers = _connected_servers()
servers = connected_servers()
valid_names = {s["name"] for s in servers}
if server_name not in valid_names:
return JSONResponse({"status": "unknown_server", "available": sorted(valid_names)})
@@ -681,8 +680,8 @@ async def mcp_meta(action: str, request: Request):
if session.sdk_session_id:
session.needs_fresh_session = True
try:
from backend.apps.agents.core.ws_manager import ws_manager as _ws
await _ws.send_to_session(parent_session_id, "agent:status", {
from backend.apps.agents.core.ws_manager import ws_manager
await ws_manager.send_to_session(parent_session_id, "agent:status", {
"session_id": parent_session_id,
"status": session.status,
"session": session.model_dump(mode="json"),
@@ -749,14 +748,14 @@ async def session_compact(session_id: str):
only sets the marker; the button is the user opting into the cost).
"""
from backend.apps.agents.agent_manager import agent_manager
from backend.apps.agents.core.ws_manager import ws_manager as _ws
from backend.apps.agents.core.ws_manager import ws_manager
session = agent_manager.sessions.get(session_id)
if not session:
return JSONResponse({"error": "session not found"}, status_code=404)
did_compact = agent_manager._maybe_compact(session, force=True)
if did_compact:
session.needs_fresh_session = True
await _ws.send_to_session(session_id, "agent:context_status", {
await ws_manager.send_to_session(session_id, "agent:context_status", {
"session_id": session_id,
"reason": "compacted_manual" if did_compact else "noop",
"compacted_through_msg_id": session.compacted_through_msg_id,
@@ -768,7 +767,7 @@ async def session_compact(session_id: str):
async def session_clear(session_id: str):
"""Wipe the session's UI history AND its SDK convo state (/clear slash cmd, Reset history button)."""
from backend.apps.agents.agent_manager import agent_manager
from backend.apps.agents.core.ws_manager import ws_manager as _ws
from backend.apps.agents.core.ws_manager import ws_manager
from backend.apps.agents.core.models import MessageBranch
session = agent_manager.sessions.get(session_id)
if not session:
@@ -784,12 +783,12 @@ async def session_clear(session_id: str):
session.branches = {"main": MessageBranch(id="main")}
session.active_branch_id = "main"
session.tool_group_meta = {}
await _ws.send_to_session(session_id, "agent:status", {
await ws_manager.send_to_session(session_id, "agent:status", {
"session_id": session_id,
"status": session.status,
"session": session.model_dump(mode="json"),
})
await _ws.send_to_session(session_id, "agent:context_status", {
await ws_manager.send_to_session(session_id, "agent:context_status", {
"session_id": session_id,
"reason": "cleared",
})
@@ -841,7 +840,7 @@ if __name__ == "__main__":
import uvicorn.config
class _ReadyServer(uvicorn.Server):
class P_ReadyServer(uvicorn.Server):
"""Subclass that prints a machine-readable READY line on startup."""
async def startup(self, sockets=None):
await super().startup(sockets)
@@ -851,6 +850,6 @@ if __name__ == "__main__":
uvicorn.run("backend.main:app", host=args.host, port=args.port, reload=True)
else:
config = uvicorn.Config("backend.main:app", host=args.host, port=args.port)
server = _ReadyServer(config)
server = P_ReadyServer(config)
import asyncio
asyncio.run(server.serve())
+6 -6
View File
@@ -12,7 +12,7 @@ import pytest
@pytest.fixture(autouse=True)
def _isolate_browser_state(monkeypatch):
def isolate_browser_state(monkeypatch):
skills_dir = tempfile.mkdtemp(prefix="os_skills_")
metrics_dir = tempfile.mkdtemp(prefix="os_metrics_")
playbook_dir = tempfile.mkdtemp(prefix="os_playbook_")
@@ -20,7 +20,7 @@ def _isolate_browser_state(monkeypatch):
monkeypatch.setenv("OPENSWARM_BROWSER_METRICS_DIR", metrics_dir)
monkeypatch.setenv("OPENSWARM_BROWSER_PLAYBOOK_DIR", playbook_dir)
def _reset():
def reset():
for mod in ("browser_skills", "browser_playbook"):
try:
m = __import__(f"backend.apps.agents.browser.{mod}", fromlist=[mod])
@@ -30,10 +30,10 @@ def _isolate_browser_state(monkeypatch):
# metrics caches its dir at first use; drop it so each test writes
# where ITS env var points, not where the first test's pointed
try:
from backend.apps.agents.browser import browser_metrics as _bm
_bm._metrics_dir_cache = None
from backend.apps.agents.browser import browser_metrics
browser_metrics.P_METRICS_DIR_CACHE = None
except Exception:
pass
_reset()
reset()
yield
_reset()
reset()
+4 -4
View File
@@ -19,13 +19,13 @@ def client():
"""Returns a TestClient pre-loaded with the local backend's auth token
so the LocalAuthMiddleware doesn't reject our requests with 401."""
import backend.auth as auth_mod
if not auth_mod._TOKEN:
if not auth_mod.TOKEN:
# Tests sometimes run without backend.main's startup hook firing.
# Generate a token directly so request_matches_token has something
# to compare against.
import secrets
auth_mod._TOKEN = secrets.token_urlsafe(32)
return TestClient(app, headers={"Authorization": f"Bearer {auth_mod._TOKEN}"})
auth_mod.TOKEN = secrets.token_urlsafe(32)
return TestClient(app, headers={"Authorization": f"Bearer {auth_mod.TOKEN}"})
@pytest.fixture
@@ -206,7 +206,7 @@ def test_dev_token_is_dev_only():
os.environ.pop("OPENSWARM_PACKAGED", None)
r = noauth.get("/api/dev/token")
assert r.status_code == 200
assert r.json()["token"] == auth_mod._TOKEN
assert r.json()["token"] == auth_mod.TOKEN
os.environ["OPENSWARM_PACKAGED"] = "1"
try: