Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 41 additions & 7 deletions backend/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@
from app.runtime.engine_manager import engine_manager
from app.runtime.hardware import check_hardware_readiness, get_gpu_stats
from app.runtime.installer import installer, mirror_manager
from app.runtime.credentials import credentials_manager
from app.runtime.credentials import credentials_manager, redact_key
from app.storage.model_store import model_store
from app.schemas.creative import CreativeActionRequest, CreativeActionResult
from app.schemas.cloud import (
Expand Down Expand Up @@ -855,7 +855,11 @@ async def list_cloud_providers() -> List[CloudProviderInfo]:
@app.post("/api/v1/cloud/credentials", response_model=CloudProviderInfo)
async def set_cloud_credential(req: SetCredentialRequest) -> CloudProviderInfo:
"""Store or update a BYOK cloud API key locally without exposure."""
credentials_manager.set_key(req.provider_id, req.api_key)
try:
credentials_manager.set_key(req.provider_id, req.api_key)
except Exception as e:
logger.error(f"Failed to persist credential: {e}")
raise HTTPException(status_code=500, detail=f"Failed to persist credentials: {e}")
providers = credentials_manager.list_providers()
target = next((p for p in providers if p.id == req.provider_id), None)
if not target:
Expand All @@ -872,7 +876,11 @@ async def test_cloud_credential(req: TestKeyRequest) -> TestKeyResult:
@app.delete("/api/v1/cloud/credentials/{provider_id}")
async def delete_cloud_credential(provider_id: CloudProviderId) -> dict[str, bool]:
"""Delete a stored BYOK API key."""
credentials_manager.delete_key(provider_id)
try:
credentials_manager.delete_key(provider_id)
except Exception as e:
logger.error(f"Failed to delete credential: {e}")
raise HTTPException(status_code=500, detail=f"Failed to update stored credentials: {e}")
return {"success": True}


Expand All @@ -883,22 +891,48 @@ async def delete_cloud_credential(provider_id: CloudProviderId) -> dict[str, boo
@app.get("/api/v1/agent/llm/config", response_model=LLMConfig)
async def get_agent_llm_config() -> LLMConfig:
"""Get active LLM provider configuration for the Agent base (llama-server / SiliconFlow / OpenAI)."""
return credentials_manager.get_llm_config()
cfg = credentials_manager.get_llm_config()
redacted_api_key = redact_key(cfg.api_key) if cfg.api_key else ""
return LLMConfig(
provider=cfg.provider,
model=cfg.model,
base_url=cfg.base_url,
api_key=redacted_api_key,
temperature=cfg.temperature,
enabled=cfg.enabled,
)


@app.post("/api/v1/agent/llm/config", response_model=LLMConfig)
async def set_agent_llm_config(req: SetLLMConfigRequest) -> LLMConfig:
"""Save or update LLM provider configuration."""
existing = credentials_manager.get_llm_config()
effective_api_key = req.api_key
if not effective_api_key or "..." in effective_api_key or effective_api_key == "****":
effective_api_key = existing.api_key

cfg = LLMConfig(
provider=req.provider,
model=req.model,
base_url=req.base_url,
api_key=req.api_key,
api_key=effective_api_key,
temperature=req.temperature,
enabled=req.enabled,
)
credentials_manager.set_llm_config(cfg)
return cfg
try:
credentials_manager.set_llm_config(cfg)
except Exception as e:
logger.error(f"Failed to persist LLM config: {e}")
raise HTTPException(status_code=500, detail=f"Failed to persist LLM configuration: {e}")

return LLMConfig(
provider=cfg.provider,
model=cfg.model,
base_url=cfg.base_url,
api_key=redact_key(cfg.api_key) if cfg.api_key else "",
temperature=cfg.temperature,
enabled=cfg.enabled,
)


@app.post("/api/v1/agent/llm/test", response_model=TestKeyResult)
Expand Down
115 changes: 112 additions & 3 deletions backend/app/runtime/credentials.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,12 @@
and strictly redacted in all diagnostic and status APIs.
"""

import base64
import json
import logging
import os
from pathlib import Path
import sys
from typing import Dict, List, Optional
import httpx

Expand All @@ -24,6 +26,68 @@
logger = logging.getLogger(__name__)


def _is_windows() -> bool:
return sys.platform == "win32"


def _protect_dpapi(data: bytes) -> bytes:
import ctypes
import ctypes.wintypes

class DATA_BLOB(ctypes.Structure):
_fields_ = [
("cbData", ctypes.wintypes.DWORD),
("pbData", ctypes.POINTER(ctypes.c_byte)),
]

blob_in = DATA_BLOB(len(data), ctypes.cast(ctypes.create_string_buffer(data), ctypes.POINTER(ctypes.c_byte)))
blob_out = DATA_BLOB()
# 0x1 = CRYPTPROTECT_UI_FORBIDDEN
if not ctypes.windll.crypt32.CryptProtectData(
ctypes.byref(blob_in),
"BerryCredential",
None,
None,
None,
0x1,
ctypes.byref(blob_out),
):
raise OSError("CryptProtectData failed to encrypt credential payload")
try:
return ctypes.string_at(blob_out.pbData, blob_out.cbData)
finally:
ctypes.windll.kernel32.LocalFree(blob_out.pbData)


def _unprotect_dpapi(data: bytes) -> bytes:
import ctypes
import ctypes.wintypes

class DATA_BLOB(ctypes.Structure):
_fields_ = [
("cbData", ctypes.wintypes.DWORD),
("pbData", ctypes.POINTER(ctypes.c_byte)),
]

blob_in = DATA_BLOB(len(data), ctypes.cast(ctypes.create_string_buffer(data), ctypes.POINTER(ctypes.c_byte)))
blob_out = DATA_BLOB()
# 0x1 = CRYPTPROTECT_UI_FORBIDDEN
if not ctypes.windll.crypt32.CryptUnprotectData(
ctypes.byref(blob_in),
None,
None,
None,
None,
0x1,
ctypes.byref(blob_out),
):
raise OSError("CryptUnprotectData failed to decrypt credential payload")
try:
return ctypes.string_at(blob_out.pbData, blob_out.cbData)
finally:
ctypes.windll.kernel32.LocalFree(blob_out.pbData)


def redact_key(key: str) -> str:
"""Mask an API key for safe display, preserving prefix and last 4 characters."""
if not key:
Expand All @@ -47,18 +111,42 @@ def __init__(self, data_dir: Optional[Path] = None) -> None:
def _load(self) -> None:
if self.creds_file.is_file():
try:
data = json.loads(self.creds_file.read_text(encoding="utf-8"))
content = self.creds_file.read_text(encoding="utf-8")
data = json.loads(content)
if isinstance(data, dict):
self._memory_creds = data
if data.get("encrypted") and data.get("format") == "dpapi" and "data" in data:
raw_bytes = base64.b64decode(data["data"])
decrypted = _unprotect_dpapi(raw_bytes).decode("utf-8")
unpacked = json.loads(decrypted)
if isinstance(unpacked, dict):
self._memory_creds = unpacked
else:
self._memory_creds = data
except Exception as e:
logger.warning(f"Error loading credentials from {self.creds_file}: {e}")

def _save(self) -> None:
self.data_dir.mkdir(parents=True, exist_ok=True)
raw_text = json.dumps(self._memory_creds, indent=2)
try:
self.creds_file.write_text(json.dumps(self._memory_creds, indent=2), encoding="utf-8")
if _is_windows():
encrypted_blob = _protect_dpapi(raw_text.encode("utf-8"))
envelope = {
"encrypted": True,
"format": "dpapi",
"data": base64.b64encode(encrypted_blob).decode("ascii"),
}
self.creds_file.write_text(json.dumps(envelope, indent=2), encoding="utf-8")
else:
self.creds_file.write_text(raw_text, encoding="utf-8")
if hasattr(os, "chmod") and os.name != "nt":
try:
os.chmod(self.creds_file, 0o600)
except Exception:
pass
except Exception as e:
logger.error(f"Error saving credentials: {e}")
raise OSError(f"Failed to persist credentials: {e}") from e

def get_key(self, provider_id: CloudProviderId) -> Optional[str]:
"""Resolve API key from stored BYOK credentials or environment variable."""
Expand Down Expand Up @@ -164,6 +252,27 @@ def set_llm_config(self, config: LLMConfig) -> None:
async def test_llm_connection(self, config: Optional[LLMConfig] = None) -> TestKeyResult:
"""Test connection to LLM provider endpoint (llama-server / OpenAI / SiliconFlow / DeepSeek)."""
cfg = config or self.get_llm_config()
if cfg.api_key and ("..." in cfg.api_key or cfg.api_key == "****"):
stored = self.get_llm_config()
cfg = LLMConfig(
provider=cfg.provider,
model=cfg.model,
base_url=cfg.base_url,
api_key=stored.api_key,
temperature=cfg.temperature,
enabled=cfg.enabled,
)
elif not cfg.api_key and config is not None:
stored = self.get_llm_config()
if stored.provider == cfg.provider and stored.api_key:
cfg = LLMConfig(
provider=cfg.provider,
model=cfg.model,
base_url=cfg.base_url,
api_key=stored.api_key,
temperature=cfg.temperature,
enabled=cfg.enabled,
)
base_url = cfg.base_url.rstrip("/")
# Target OpenAI-compatible /models endpoint
models_url = f"{base_url}/models" if not base_url.endswith("/models") else base_url
Expand Down
140 changes: 140 additions & 0 deletions backend/tests/test_credentials_security.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
import json
from pathlib import Path
import tempfile
from unittest.mock import patch
import pytest
from fastapi.testclient import TestClient

from app.main import app, credentials_manager
from app.runtime.credentials import CredentialManager, redact_key, _is_windows
from app.schemas.cloud import CloudProviderId, LLMConfig, SetLLMConfigRequest

client = TestClient(app)


def test_redact_key_helper():
assert redact_key("") == ""
assert redact_key("12345") == "****"
assert redact_key("12345678") == "****"
assert redact_key("sk-1234567890abcdef") == "sk-...cdef"


def test_credential_encryption_on_disk():
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
mgr = CredentialManager(data_dir=tmp_path)
dummy_secret = "super-confidential-secret-key-xyz-987"
mgr.set_key(CloudProviderId.OPENAI, dummy_secret)

raw_file_content = mgr.creds_file.read_text(encoding="utf-8")
if _is_windows():
# In Windows, verify DPAPI encrypted envelope is used and raw secret is absent
assert dummy_secret not in raw_file_content
data = json.loads(raw_file_content)
assert data.get("encrypted") is True
assert data.get("format") == "dpapi"
assert "data" in data

# Reloading in another manager instance must recover the original secret
mgr_reloaded = CredentialManager(data_dir=tmp_path)
assert mgr_reloaded.get_key(CloudProviderId.OPENAI) == dummy_secret


def test_credential_legacy_plaintext_backward_compatibility():
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
creds_file = tmp_path / "credentials.json"
legacy_data = {
"openai": "sk-legacy-unencrypted-key",
"fal": "fal-legacy-key",
}
creds_file.write_text(json.dumps(legacy_data), encoding="utf-8")

mgr = CredentialManager(data_dir=tmp_path)
assert mgr.get_key(CloudProviderId.OPENAI) == "sk-legacy-unencrypted-key"
assert mgr.get_key(CloudProviderId.FAL) == "fal-legacy-key"


def test_credential_save_failure_raises_error():
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
mgr = CredentialManager(data_dir=tmp_path)
with patch.object(Path, "write_text", side_effect=OSError("Disk full")):
with pytest.raises(OSError, match="Failed to persist credentials"):
mgr.set_key(CloudProviderId.OPENAI, "sk-test")


def test_llm_config_endpoints_redact_secrets():
dummy_key = "sk-llm-secret-token-abcdef123456"
stored_cfg = LLMConfig(
provider="openai",
model="gpt-4o",
base_url="https://api.openai.com/v1",
api_key=dummy_key,
temperature=0.7,
enabled=True,
)

with patch.object(credentials_manager, "get_llm_config", return_value=stored_cfg):
resp = client.get("/api/v1/agent/llm/config")
assert resp.status_code == 200
data = resp.json()
assert data["api_key"] != dummy_key
assert data["api_key"] == redact_key(dummy_key)


def test_llm_config_set_preserves_existing_secret_when_redacted_or_empty():
dummy_key = "sk-llm-secret-token-abcdef123456"
existing_cfg = LLMConfig(
provider="openai",
model="gpt-4o",
base_url="https://api.openai.com/v1",
api_key=dummy_key,
temperature=0.7,
enabled=True,
)

saved_configs = []

def mock_set_llm_config(cfg):
saved_configs.append(cfg)

with patch.object(credentials_manager, "get_llm_config", return_value=existing_cfg), \
patch.object(credentials_manager, "set_llm_config", side_effect=mock_set_llm_config):
# Case 1: user sends redacted key back
req_payload = {
"provider": "openai",
"model": "gpt-4o-mini",
"base_url": "https://api.openai.com/v1",
"api_key": redact_key(dummy_key),
"temperature": 0.5,
"enabled": True,
}
resp = client.post("/api/v1/agent/llm/config", json=req_payload)
assert resp.status_code == 200
assert len(saved_configs) == 1
assert saved_configs[0].api_key == dummy_key
assert saved_configs[0].model == "gpt-4o-mini"
assert resp.json()["api_key"] == redact_key(dummy_key)

# Case 2: user sends empty string key
req_payload["api_key"] = ""
resp = client.post("/api/v1/agent/llm/config", json=req_payload)
assert resp.status_code == 200
assert len(saved_configs) == 2
assert saved_configs[1].api_key == dummy_key


def test_llm_config_set_persistence_error_returns_500():
with patch.object(credentials_manager, "set_llm_config", side_effect=OSError("Read-only filesystem")):
req_payload = {
"provider": "openai",
"model": "gpt-4o",
"base_url": "https://api.openai.com/v1",
"api_key": "sk-new-key",
"temperature": 0.7,
"enabled": True,
}
resp = client.post("/api/v1/agent/llm/config", json=req_payload)
assert resp.status_code == 500
assert "Failed to persist LLM configuration" in resp.json()["detail"]
Loading