iis-access 0.1.1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- iis_access-0.1.1/.gitignore +13 -0
- iis_access-0.1.1/PKG-INFO +48 -0
- iis_access-0.1.1/README.md +22 -0
- iis_access-0.1.1/pyproject.toml +44 -0
- iis_access-0.1.1/src/iis_access/__init__.py +19 -0
- iis_access-0.1.1/src/iis_access/_types.py +23 -0
- iis_access-0.1.1/src/iis_access/auth.py +150 -0
- iis_access-0.1.1/src/iis_access/config.py +45 -0
- iis_access-0.1.1/src/iis_access/managed_keys.py +54 -0
- iis_access-0.1.1/src/iis_access/nexus.py +19 -0
- iis_access-0.1.1/src/iis_access/provider.py +142 -0
- iis_access-0.1.1/src/iis_access/sync.py +30 -0
- iis_access-0.1.1/src/iis_access/usage.py +40 -0
- iis_access-0.1.1/tests/__init__.py +0 -0
- iis_access-0.1.1/tests/conftest.py +14 -0
- iis_access-0.1.1/tests/test_auth.py +120 -0
- iis_access-0.1.1/tests/test_graceful_degradation.py +96 -0
- iis_access-0.1.1/tests/test_managed_keys.py +115 -0
- iis_access-0.1.1/tests/test_provider.py +142 -0
- iis_access-0.1.1/tests/test_usage.py +73 -0
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: iis-access
|
|
3
|
+
Version: 0.1.1
|
|
4
|
+
Summary: Python auth library for IIS tools ecosystem
|
|
5
|
+
Project-URL: Repository, https://github.com/hschmied/account
|
|
6
|
+
Project-URL: Homepage, https://iis.tools
|
|
7
|
+
Author-email: IIS Labs <dev@iis.tools>
|
|
8
|
+
License-Expression: MIT
|
|
9
|
+
Keywords: api-key,auth,iis,provider-resolution
|
|
10
|
+
Classifier: Development Status :: 3 - Alpha
|
|
11
|
+
Classifier: Intended Audience :: Developers
|
|
12
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
13
|
+
Classifier: Programming Language :: Python :: 3
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
17
|
+
Requires-Python: >=3.11
|
|
18
|
+
Requires-Dist: aiohttp>=3.9
|
|
19
|
+
Provides-Extra: dev
|
|
20
|
+
Requires-Dist: aioresponses>=0.7; extra == 'dev'
|
|
21
|
+
Requires-Dist: pytest-asyncio>=0.24; extra == 'dev'
|
|
22
|
+
Requires-Dist: pytest>=8.0; extra == 'dev'
|
|
23
|
+
Provides-Extra: keyring
|
|
24
|
+
Requires-Dist: keyring>=25.0; extra == 'keyring'
|
|
25
|
+
Description-Content-Type: text/markdown
|
|
26
|
+
|
|
27
|
+
# iis-access
|
|
28
|
+
|
|
29
|
+
Python auth library for IIS tools ecosystem. Provides optional authentication, provider resolution, managed key retrieval, and usage reporting for Python-based IIS CLI tools (aria, sonar, iris).
|
|
30
|
+
|
|
31
|
+
## Installation
|
|
32
|
+
|
|
33
|
+
```bash
|
|
34
|
+
pip install iis-access
|
|
35
|
+
```
|
|
36
|
+
|
|
37
|
+
## Quick Start
|
|
38
|
+
|
|
39
|
+
```python
|
|
40
|
+
from iis_access import resolve_provider, report_usage
|
|
41
|
+
|
|
42
|
+
# Resolve the best available provider for a task
|
|
43
|
+
resolution = await resolve_provider("tts", preferred="elevenlabs")
|
|
44
|
+
print(resolution.provider, resolution.method)
|
|
45
|
+
|
|
46
|
+
# Report usage (fire-and-forget, never raises)
|
|
47
|
+
await report_usage("aria", credits=1.5)
|
|
48
|
+
```
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
# iis-access
|
|
2
|
+
|
|
3
|
+
Python auth library for IIS tools ecosystem. Provides optional authentication, provider resolution, managed key retrieval, and usage reporting for Python-based IIS CLI tools (aria, sonar, iris).
|
|
4
|
+
|
|
5
|
+
## Installation
|
|
6
|
+
|
|
7
|
+
```bash
|
|
8
|
+
pip install iis-access
|
|
9
|
+
```
|
|
10
|
+
|
|
11
|
+
## Quick Start
|
|
12
|
+
|
|
13
|
+
```python
|
|
14
|
+
from iis_access import resolve_provider, report_usage
|
|
15
|
+
|
|
16
|
+
# Resolve the best available provider for a task
|
|
17
|
+
resolution = await resolve_provider("tts", preferred="elevenlabs")
|
|
18
|
+
print(resolution.provider, resolution.method)
|
|
19
|
+
|
|
20
|
+
# Report usage (fire-and-forget, never raises)
|
|
21
|
+
await report_usage("aria", credits=1.5)
|
|
22
|
+
```
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["hatchling"]
|
|
3
|
+
build-backend = "hatchling.build"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "iis-access"
|
|
7
|
+
version = "0.1.1"
|
|
8
|
+
description = "Python auth library for IIS tools ecosystem"
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
requires-python = ">=3.11"
|
|
11
|
+
license = "MIT"
|
|
12
|
+
authors = [{name = "IIS Labs", email = "dev@iis.tools"}]
|
|
13
|
+
keywords = ["iis", "auth", "api-key", "provider-resolution"]
|
|
14
|
+
dependencies = [
|
|
15
|
+
"aiohttp>=3.9",
|
|
16
|
+
]
|
|
17
|
+
classifiers = [
|
|
18
|
+
"Development Status :: 3 - Alpha",
|
|
19
|
+
"Intended Audience :: Developers",
|
|
20
|
+
"License :: OSI Approved :: MIT License",
|
|
21
|
+
"Programming Language :: Python :: 3",
|
|
22
|
+
"Programming Language :: Python :: 3.11",
|
|
23
|
+
"Programming Language :: Python :: 3.12",
|
|
24
|
+
"Programming Language :: Python :: 3.13",
|
|
25
|
+
]
|
|
26
|
+
|
|
27
|
+
[project.urls]
|
|
28
|
+
Repository = "https://github.com/hschmied/account"
|
|
29
|
+
Homepage = "https://iis.tools"
|
|
30
|
+
|
|
31
|
+
[project.optional-dependencies]
|
|
32
|
+
keyring = ["keyring>=25.0"]
|
|
33
|
+
dev = [
|
|
34
|
+
"pytest>=8.0",
|
|
35
|
+
"pytest-asyncio>=0.24",
|
|
36
|
+
"aioresponses>=0.7",
|
|
37
|
+
]
|
|
38
|
+
|
|
39
|
+
[tool.hatch.build.targets.wheel]
|
|
40
|
+
packages = ["src/iis_access"]
|
|
41
|
+
|
|
42
|
+
[tool.pytest.ini_options]
|
|
43
|
+
asyncio_mode = "auto"
|
|
44
|
+
testpaths = ["tests"]
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
"""IIS Tools Python access library - optional auth integration for CLI tools."""
|
|
2
|
+
|
|
3
|
+
from ._types import AuthContext, ProviderResolution, NoProviderAvailable
|
|
4
|
+
from .auth import get_auth_context, is_authenticated
|
|
5
|
+
from .provider import resolve_provider, register_free_provider
|
|
6
|
+
from .managed_keys import get_managed_key
|
|
7
|
+
from .usage import report_usage
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"AuthContext",
|
|
11
|
+
"ProviderResolution",
|
|
12
|
+
"NoProviderAvailable",
|
|
13
|
+
"get_auth_context",
|
|
14
|
+
"is_authenticated",
|
|
15
|
+
"resolve_provider",
|
|
16
|
+
"register_free_provider",
|
|
17
|
+
"get_managed_key",
|
|
18
|
+
"report_usage",
|
|
19
|
+
]
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
@dataclass
|
|
5
|
+
class AuthContext:
|
|
6
|
+
account_id: str
|
|
7
|
+
access_token: str
|
|
8
|
+
plan: str # "free", "pro", "enterprise"
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@dataclass
|
|
12
|
+
class ProviderResolution:
|
|
13
|
+
provider: str # "openai", "elevenlabs", "edge-tts", etc.
|
|
14
|
+
method: str # "nexus", "managed_key", "local_key", "free"
|
|
15
|
+
api_key: str | None # The key to use (None for free providers)
|
|
16
|
+
endpoint: str | None # Override endpoint (for Nexus)
|
|
17
|
+
credits_estimated: float | None # Estimated credit cost
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class NoProviderAvailable(Exception):
|
|
21
|
+
"""Raised when no provider can fulfill the request."""
|
|
22
|
+
|
|
23
|
+
pass
|
|
@@ -0,0 +1,150 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import logging
|
|
3
|
+
|
|
4
|
+
from ._types import AuthContext
|
|
5
|
+
from .config import get_account_url, get_credentials, is_auth_disabled
|
|
6
|
+
|
|
7
|
+
logger = logging.getLogger("iis_access")
|
|
8
|
+
|
|
9
|
+
_cached_context: AuthContext | None = None
|
|
10
|
+
_refresh_lock: asyncio.Lock | None = None
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _get_lock() -> asyncio.Lock:
|
|
14
|
+
global _refresh_lock
|
|
15
|
+
if _refresh_lock is None:
|
|
16
|
+
_refresh_lock = asyncio.Lock()
|
|
17
|
+
return _refresh_lock
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
async def get_auth_context() -> AuthContext | None:
|
|
21
|
+
"""Get authentication context. Returns None if not authenticated."""
|
|
22
|
+
global _cached_context
|
|
23
|
+
|
|
24
|
+
if is_auth_disabled():
|
|
25
|
+
return None
|
|
26
|
+
|
|
27
|
+
if _cached_context is not None:
|
|
28
|
+
return _cached_context
|
|
29
|
+
|
|
30
|
+
async with _get_lock():
|
|
31
|
+
# Double-check after acquiring lock
|
|
32
|
+
if _cached_context is not None:
|
|
33
|
+
return _cached_context
|
|
34
|
+
|
|
35
|
+
creds = get_credentials()
|
|
36
|
+
if not creds or "accountId" not in creds:
|
|
37
|
+
return None
|
|
38
|
+
|
|
39
|
+
account_id = creds["accountId"]
|
|
40
|
+
refresh_token = _get_refresh_token(account_id)
|
|
41
|
+
if not refresh_token:
|
|
42
|
+
return None
|
|
43
|
+
|
|
44
|
+
# Exchange refresh token for access token
|
|
45
|
+
try:
|
|
46
|
+
import aiohttp
|
|
47
|
+
|
|
48
|
+
account_url = get_account_url()
|
|
49
|
+
timeout = aiohttp.ClientTimeout(total=5)
|
|
50
|
+
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
51
|
+
resp = await session.post(
|
|
52
|
+
f"{account_url}/auth/refresh",
|
|
53
|
+
json={"refreshToken": refresh_token},
|
|
54
|
+
)
|
|
55
|
+
if resp.status != 200:
|
|
56
|
+
logger.info(
|
|
57
|
+
"Token refresh failed (status %d), continuing without auth",
|
|
58
|
+
resp.status,
|
|
59
|
+
)
|
|
60
|
+
return None
|
|
61
|
+
|
|
62
|
+
data = await resp.json()
|
|
63
|
+
access_token = data.get("accessToken")
|
|
64
|
+
new_refresh_token = data.get("refreshToken")
|
|
65
|
+
plan = data.get("plan") or _extract_plan_from_jwt(access_token)
|
|
66
|
+
|
|
67
|
+
if not access_token:
|
|
68
|
+
return None
|
|
69
|
+
|
|
70
|
+
# Store new refresh token back
|
|
71
|
+
if new_refresh_token:
|
|
72
|
+
_store_refresh_token(account_id, new_refresh_token)
|
|
73
|
+
|
|
74
|
+
_cached_context = AuthContext(
|
|
75
|
+
account_id=account_id,
|
|
76
|
+
access_token=access_token,
|
|
77
|
+
plan=plan,
|
|
78
|
+
)
|
|
79
|
+
return _cached_context
|
|
80
|
+
|
|
81
|
+
except Exception as e:
|
|
82
|
+
logger.warning("Auth failed (non-fatal): %s", e)
|
|
83
|
+
return None
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def is_authenticated() -> bool:
|
|
87
|
+
"""Synchronous check if there are credentials available."""
|
|
88
|
+
creds = get_credentials()
|
|
89
|
+
return creds is not None and "accountId" in creds
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def _get_refresh_token(account_id: str) -> str | None:
|
|
93
|
+
"""Try keyring first, then fall back to credentials file."""
|
|
94
|
+
# Try keyring (optional dependency)
|
|
95
|
+
kr = _get_keyring()
|
|
96
|
+
if kr:
|
|
97
|
+
try:
|
|
98
|
+
token = kr.get_password("iis.tools", account_id)
|
|
99
|
+
if token:
|
|
100
|
+
return token
|
|
101
|
+
except Exception:
|
|
102
|
+
pass
|
|
103
|
+
|
|
104
|
+
# Fall back to credentials file
|
|
105
|
+
creds = get_credentials()
|
|
106
|
+
if creds and "refreshToken" in creds:
|
|
107
|
+
return creds["refreshToken"]
|
|
108
|
+
|
|
109
|
+
return None
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def _store_refresh_token(account_id: str, token: str) -> None:
|
|
113
|
+
"""Store refresh token in keyring if available."""
|
|
114
|
+
kr = _get_keyring()
|
|
115
|
+
if kr:
|
|
116
|
+
try:
|
|
117
|
+
kr.set_password("iis.tools", account_id, token)
|
|
118
|
+
except Exception as e:
|
|
119
|
+
logger.debug("Failed to store refresh token in keyring: %s", e)
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def _extract_plan_from_jwt(token: str | None) -> str:
|
|
123
|
+
"""Extract plan claim from JWT payload without verifying signature."""
|
|
124
|
+
if not token:
|
|
125
|
+
return "free"
|
|
126
|
+
try:
|
|
127
|
+
import base64, json as _json
|
|
128
|
+
payload = token.split(".")[1]
|
|
129
|
+
payload += "=" * (4 - len(payload) % 4)
|
|
130
|
+
claims = _json.loads(base64.urlsafe_b64decode(payload))
|
|
131
|
+
return claims.get("plan", "free")
|
|
132
|
+
except Exception:
|
|
133
|
+
return "free"
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _get_keyring():
|
|
137
|
+
"""Lazy import keyring (optional dependency)."""
|
|
138
|
+
try:
|
|
139
|
+
import keyring
|
|
140
|
+
|
|
141
|
+
return keyring
|
|
142
|
+
except ImportError:
|
|
143
|
+
return None
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _reset_auth_cache() -> None:
|
|
147
|
+
"""Reset auth cache (for testing)."""
|
|
148
|
+
global _cached_context, _refresh_lock
|
|
149
|
+
_cached_context = None
|
|
150
|
+
_refresh_lock = None
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
_IIS_DIR = Path.home() / ".iis"
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def get_account_url() -> str:
|
|
9
|
+
if url := os.environ.get("IIS_ACCOUNT_URL"):
|
|
10
|
+
return url.rstrip("/")
|
|
11
|
+
config = _load_config()
|
|
12
|
+
return config.get("accountUrl", "https://account.iis.tools").rstrip("/")
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def get_nexus_url() -> str:
|
|
16
|
+
if url := os.environ.get("IIS_NEXUS_URL"):
|
|
17
|
+
return url.rstrip("/")
|
|
18
|
+
config = _load_config()
|
|
19
|
+
return config.get("nexusUrl", "http://localhost:9666").rstrip("/")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def is_auth_disabled() -> bool:
|
|
23
|
+
return os.environ.get("IIS_NO_AUTH") == "1"
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def get_credentials() -> dict | None:
|
|
27
|
+
path = _IIS_DIR / "credentials.json"
|
|
28
|
+
if not path.exists():
|
|
29
|
+
return None
|
|
30
|
+
try:
|
|
31
|
+
with open(path) as f:
|
|
32
|
+
return json.load(f)
|
|
33
|
+
except (json.JSONDecodeError, OSError):
|
|
34
|
+
return None
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def _load_config() -> dict:
|
|
38
|
+
path = _IIS_DIR / "config.json"
|
|
39
|
+
if not path.exists():
|
|
40
|
+
return {}
|
|
41
|
+
try:
|
|
42
|
+
with open(path) as f:
|
|
43
|
+
return json.load(f)
|
|
44
|
+
except (json.JSONDecodeError, OSError):
|
|
45
|
+
return {}
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from .auth import get_auth_context
|
|
4
|
+
from .config import get_account_url
|
|
5
|
+
|
|
6
|
+
logger = logging.getLogger("iis_access")
|
|
7
|
+
|
|
8
|
+
# In-memory cache for managed keys (process lifetime)
|
|
9
|
+
_key_cache: dict[str, str] = {}
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
async def get_managed_key(provider: str) -> str | None:
|
|
13
|
+
"""
|
|
14
|
+
Retrieve user's managed API key for a provider from iis-account.
|
|
15
|
+
Returns None if not authenticated, no key exists, or unreachable.
|
|
16
|
+
"""
|
|
17
|
+
if provider in _key_cache:
|
|
18
|
+
return _key_cache[provider]
|
|
19
|
+
|
|
20
|
+
context = await get_auth_context()
|
|
21
|
+
if context is None:
|
|
22
|
+
return None
|
|
23
|
+
|
|
24
|
+
try:
|
|
25
|
+
import aiohttp
|
|
26
|
+
|
|
27
|
+
account_url = get_account_url()
|
|
28
|
+
timeout = aiohttp.ClientTimeout(total=5)
|
|
29
|
+
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
30
|
+
resp = await session.get(
|
|
31
|
+
f"{account_url}/keys/managed/{provider}",
|
|
32
|
+
headers={"Authorization": f"Bearer {context.access_token}"},
|
|
33
|
+
)
|
|
34
|
+
if resp.status == 200:
|
|
35
|
+
data = await resp.json()
|
|
36
|
+
key = data.get("key")
|
|
37
|
+
if key:
|
|
38
|
+
_key_cache[provider] = key
|
|
39
|
+
return key
|
|
40
|
+
elif resp.status == 404:
|
|
41
|
+
logger.debug("No managed key for provider '%s'", provider)
|
|
42
|
+
else:
|
|
43
|
+
logger.debug(
|
|
44
|
+
"Managed key retrieval failed (status %d)", resp.status
|
|
45
|
+
)
|
|
46
|
+
except Exception as e:
|
|
47
|
+
logger.debug("Managed key retrieval failed (non-fatal): %s", e)
|
|
48
|
+
|
|
49
|
+
return None
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _reset_key_cache() -> None:
|
|
53
|
+
"""Reset managed key cache (for testing)."""
|
|
54
|
+
_key_cache.clear()
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from .config import get_nexus_url
|
|
4
|
+
|
|
5
|
+
logger = logging.getLogger("iis_access")
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
async def is_nexus_running() -> bool:
|
|
9
|
+
"""Check if Nexus local agent is running and responsive."""
|
|
10
|
+
try:
|
|
11
|
+
import aiohttp
|
|
12
|
+
|
|
13
|
+
nexus_url = get_nexus_url()
|
|
14
|
+
timeout = aiohttp.ClientTimeout(total=1)
|
|
15
|
+
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
16
|
+
resp = await session.get(f"{nexus_url}/health")
|
|
17
|
+
return resp.status == 200
|
|
18
|
+
except Exception:
|
|
19
|
+
return False
|
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import logging
|
|
3
|
+
import shutil
|
|
4
|
+
from typing import Callable
|
|
5
|
+
|
|
6
|
+
from ._types import ProviderResolution, NoProviderAvailable
|
|
7
|
+
from .nexus import is_nexus_running
|
|
8
|
+
from .managed_keys import get_managed_key
|
|
9
|
+
from .auth import get_auth_context
|
|
10
|
+
|
|
11
|
+
logger = logging.getLogger("iis_access")
|
|
12
|
+
|
|
13
|
+
# Known env var names for local API keys
|
|
14
|
+
_ENV_KEY_MAP: dict[str, str] = {
|
|
15
|
+
"openai": "OPENAI_API_KEY",
|
|
16
|
+
"anthropic": "ANTHROPIC_API_KEY",
|
|
17
|
+
"elevenlabs": "ELEVENLABS_API_KEY",
|
|
18
|
+
"azure-openai": "AZURE_OPENAI_API_KEY",
|
|
19
|
+
"google-ai": "GOOGLE_AI_API_KEY",
|
|
20
|
+
"mistral": "MISTRAL_API_KEY",
|
|
21
|
+
"groq": "GROQ_API_KEY",
|
|
22
|
+
"deepseek": "DEEPSEEK_API_KEY",
|
|
23
|
+
}
|
|
24
|
+
|
|
25
|
+
# Free provider registry
|
|
26
|
+
_FREE_PROVIDERS: dict[str, list[tuple[str, Callable[[], bool] | None]]] = {
|
|
27
|
+
"tts": [
|
|
28
|
+
("edge-tts", None), # Always available (pip package)
|
|
29
|
+
("mac-say", lambda: shutil.which("say") is not None),
|
|
30
|
+
],
|
|
31
|
+
"transcription": [
|
|
32
|
+
("whisper-local", lambda: shutil.which("whisper") is not None),
|
|
33
|
+
],
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def register_free_provider(
|
|
38
|
+
task_type: str,
|
|
39
|
+
provider: str,
|
|
40
|
+
check: Callable[[], bool] | None = None,
|
|
41
|
+
) -> None:
|
|
42
|
+
"""Register a free provider for a task type."""
|
|
43
|
+
if task_type not in _FREE_PROVIDERS:
|
|
44
|
+
_FREE_PROVIDERS[task_type] = []
|
|
45
|
+
_FREE_PROVIDERS[task_type].append((provider, check))
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
async def resolve_provider(
|
|
49
|
+
task_type: str,
|
|
50
|
+
preferred: str | None = None,
|
|
51
|
+
require_capability: str | None = None,
|
|
52
|
+
) -> ProviderResolution:
|
|
53
|
+
"""
|
|
54
|
+
Resolve the best provider for the given task type.
|
|
55
|
+
|
|
56
|
+
Resolution order:
|
|
57
|
+
1. Nexus with pooled credits
|
|
58
|
+
2. Nexus with managed key
|
|
59
|
+
3. Local API key (env var)
|
|
60
|
+
4. Free provider
|
|
61
|
+
5. Raise NoProviderAvailable
|
|
62
|
+
"""
|
|
63
|
+
context = await get_auth_context()
|
|
64
|
+
|
|
65
|
+
# Branch 1 & 2: Nexus
|
|
66
|
+
if await is_nexus_running():
|
|
67
|
+
nexus_url = _get_nexus_endpoint()
|
|
68
|
+
|
|
69
|
+
# Branch 1: Nexus with pooled credits (if authenticated with plan)
|
|
70
|
+
if context and context.plan in ("pro", "enterprise"):
|
|
71
|
+
return ProviderResolution(
|
|
72
|
+
provider=preferred or "nexus-default",
|
|
73
|
+
method="nexus",
|
|
74
|
+
api_key=None,
|
|
75
|
+
endpoint=nexus_url,
|
|
76
|
+
credits_estimated=None, # Nexus handles pricing
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
# Branch 2: Nexus with managed key
|
|
80
|
+
if context and preferred:
|
|
81
|
+
managed_key = await get_managed_key(preferred)
|
|
82
|
+
if managed_key:
|
|
83
|
+
return ProviderResolution(
|
|
84
|
+
provider=preferred,
|
|
85
|
+
method="managed_key",
|
|
86
|
+
api_key=managed_key,
|
|
87
|
+
endpoint=nexus_url,
|
|
88
|
+
credits_estimated=None,
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
# Branch 3: Local API key
|
|
92
|
+
providers_to_check = [preferred] if preferred else list(_ENV_KEY_MAP.keys())
|
|
93
|
+
for provider in providers_to_check:
|
|
94
|
+
if not provider:
|
|
95
|
+
continue
|
|
96
|
+
env_var = _ENV_KEY_MAP.get(provider)
|
|
97
|
+
if env_var:
|
|
98
|
+
key = os.environ.get(env_var)
|
|
99
|
+
if key:
|
|
100
|
+
return ProviderResolution(
|
|
101
|
+
provider=provider,
|
|
102
|
+
method="local_key",
|
|
103
|
+
api_key=key,
|
|
104
|
+
endpoint=None,
|
|
105
|
+
credits_estimated=None,
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
# Also check managed key without Nexus (direct access)
|
|
109
|
+
if context and preferred:
|
|
110
|
+
managed_key = await get_managed_key(preferred)
|
|
111
|
+
if managed_key:
|
|
112
|
+
return ProviderResolution(
|
|
113
|
+
provider=preferred,
|
|
114
|
+
method="managed_key",
|
|
115
|
+
api_key=managed_key,
|
|
116
|
+
endpoint=None,
|
|
117
|
+
credits_estimated=None,
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
# Branch 4: Free provider
|
|
121
|
+
free_providers = _FREE_PROVIDERS.get(task_type, [])
|
|
122
|
+
for name, check_fn in free_providers:
|
|
123
|
+
if check_fn is None or check_fn():
|
|
124
|
+
return ProviderResolution(
|
|
125
|
+
provider=name,
|
|
126
|
+
method="free",
|
|
127
|
+
api_key=None,
|
|
128
|
+
endpoint=None,
|
|
129
|
+
credits_estimated=0,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
# Branch 5: No provider available
|
|
133
|
+
raise NoProviderAvailable(
|
|
134
|
+
f"No provider available for task type '{task_type}'"
|
|
135
|
+
+ (f" (preferred: {preferred})" if preferred else "")
|
|
136
|
+
)
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def _get_nexus_endpoint() -> str:
|
|
140
|
+
from .config import get_nexus_url
|
|
141
|
+
|
|
142
|
+
return get_nexus_url()
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
|
|
3
|
+
from ._types import ProviderResolution
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def resolve_provider_sync(
|
|
7
|
+
task_type: str,
|
|
8
|
+
preferred: str | None = None,
|
|
9
|
+
require_capability: str | None = None,
|
|
10
|
+
) -> ProviderResolution:
|
|
11
|
+
"""Sync wrapper for resolve_provider."""
|
|
12
|
+
from .provider import resolve_provider
|
|
13
|
+
|
|
14
|
+
return asyncio.run(resolve_provider(task_type, preferred, require_capability))
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def get_managed_key_sync(provider: str) -> str | None:
|
|
18
|
+
"""Sync wrapper for get_managed_key."""
|
|
19
|
+
from .managed_keys import get_managed_key
|
|
20
|
+
|
|
21
|
+
return asyncio.run(get_managed_key(provider))
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def report_usage_sync(
|
|
25
|
+
service: str, credits: float = 0, requests: int = 1
|
|
26
|
+
) -> None:
|
|
27
|
+
"""Sync wrapper for report_usage."""
|
|
28
|
+
from .usage import report_usage
|
|
29
|
+
|
|
30
|
+
asyncio.run(report_usage(service, credits, requests))
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from .auth import get_auth_context
|
|
4
|
+
from .config import get_account_url
|
|
5
|
+
|
|
6
|
+
logger = logging.getLogger("iis_access")
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
async def report_usage(
|
|
10
|
+
service: str,
|
|
11
|
+
credits: float = 0,
|
|
12
|
+
requests: int = 1,
|
|
13
|
+
) -> None:
|
|
14
|
+
"""
|
|
15
|
+
Fire-and-forget usage report to iis-account.
|
|
16
|
+
Never raises - logs warnings on failure.
|
|
17
|
+
Skipped if not authenticated.
|
|
18
|
+
"""
|
|
19
|
+
context = await get_auth_context()
|
|
20
|
+
if context is None:
|
|
21
|
+
return
|
|
22
|
+
|
|
23
|
+
try:
|
|
24
|
+
import aiohttp
|
|
25
|
+
|
|
26
|
+
account_url = get_account_url()
|
|
27
|
+
timeout = aiohttp.ClientTimeout(total=5)
|
|
28
|
+
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
29
|
+
await session.post(
|
|
30
|
+
f"{account_url}/usage/record",
|
|
31
|
+
json={
|
|
32
|
+
"accountId": context.account_id,
|
|
33
|
+
"service": service,
|
|
34
|
+
"requests": requests,
|
|
35
|
+
"credits": credits,
|
|
36
|
+
},
|
|
37
|
+
headers={"Authorization": f"Bearer {context.access_token}"},
|
|
38
|
+
)
|
|
39
|
+
except Exception as e:
|
|
40
|
+
logger.debug("Usage reporting failed (non-fatal): %s", e)
|
|
File without changes
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
|
|
3
|
+
from iis_access.auth import _reset_auth_cache
|
|
4
|
+
from iis_access.managed_keys import _reset_key_cache
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@pytest.fixture(autouse=True)
|
|
8
|
+
def reset_caches():
|
|
9
|
+
"""Reset all caches between tests."""
|
|
10
|
+
_reset_auth_cache()
|
|
11
|
+
_reset_key_cache()
|
|
12
|
+
yield
|
|
13
|
+
_reset_auth_cache()
|
|
14
|
+
_reset_key_cache()
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from unittest.mock import patch
|
|
3
|
+
|
|
4
|
+
import pytest
|
|
5
|
+
from aioresponses import aioresponses
|
|
6
|
+
|
|
7
|
+
from iis_access._types import AuthContext
|
|
8
|
+
from iis_access.auth import get_auth_context, is_authenticated
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@pytest.fixture
|
|
12
|
+
def mock_credentials(tmp_path):
|
|
13
|
+
"""Create a temporary credentials file."""
|
|
14
|
+
iis_dir = tmp_path / ".iis"
|
|
15
|
+
iis_dir.mkdir()
|
|
16
|
+
creds_file = iis_dir / "credentials.json"
|
|
17
|
+
|
|
18
|
+
def _write(data: dict):
|
|
19
|
+
creds_file.write_text(json.dumps(data))
|
|
20
|
+
return iis_dir
|
|
21
|
+
|
|
22
|
+
return _write
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class TestGetAuthContext:
|
|
26
|
+
async def test_returns_none_when_no_credentials(self):
|
|
27
|
+
with patch("iis_access.auth.get_credentials", return_value=None):
|
|
28
|
+
result = await get_auth_context()
|
|
29
|
+
assert result is None
|
|
30
|
+
|
|
31
|
+
async def test_returns_none_when_auth_disabled(self):
|
|
32
|
+
with patch.dict("os.environ", {"IIS_NO_AUTH": "1"}):
|
|
33
|
+
result = await get_auth_context()
|
|
34
|
+
assert result is None
|
|
35
|
+
|
|
36
|
+
async def test_returns_none_when_no_account_id(self):
|
|
37
|
+
with patch(
|
|
38
|
+
"iis_access.auth.get_credentials",
|
|
39
|
+
return_value={"refreshToken": "tok"},
|
|
40
|
+
):
|
|
41
|
+
result = await get_auth_context()
|
|
42
|
+
assert result is None
|
|
43
|
+
|
|
44
|
+
async def test_returns_none_when_no_refresh_token(self):
|
|
45
|
+
with patch(
|
|
46
|
+
"iis_access.auth.get_credentials",
|
|
47
|
+
return_value={"accountId": "acc-123"},
|
|
48
|
+
):
|
|
49
|
+
with patch("iis_access.auth._get_keyring", return_value=None):
|
|
50
|
+
result = await get_auth_context()
|
|
51
|
+
assert result is None
|
|
52
|
+
|
|
53
|
+
async def test_returns_auth_context_on_successful_refresh(self):
|
|
54
|
+
creds = {"accountId": "acc-123", "refreshToken": "iis_rt_test"}
|
|
55
|
+
|
|
56
|
+
with (
|
|
57
|
+
patch("iis_access.auth.get_credentials", return_value=creds),
|
|
58
|
+
patch("iis_access.auth._get_keyring", return_value=None),
|
|
59
|
+
patch(
|
|
60
|
+
"iis_access.auth.get_account_url",
|
|
61
|
+
return_value="http://localhost:9060",
|
|
62
|
+
),
|
|
63
|
+
aioresponses() as mocked,
|
|
64
|
+
):
|
|
65
|
+
mocked.post(
|
|
66
|
+
"http://localhost:9060/auth/refresh",
|
|
67
|
+
payload={
|
|
68
|
+
"accessToken": "jwt-token-abc",
|
|
69
|
+
"refreshToken": "iis_rt_new",
|
|
70
|
+
"plan": "pro",
|
|
71
|
+
},
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
result = await get_auth_context()
|
|
75
|
+
|
|
76
|
+
assert result is not None
|
|
77
|
+
assert isinstance(result, AuthContext)
|
|
78
|
+
assert result.account_id == "acc-123"
|
|
79
|
+
assert result.access_token == "jwt-token-abc"
|
|
80
|
+
assert result.plan == "pro"
|
|
81
|
+
|
|
82
|
+
async def test_returns_none_on_refresh_failure(self):
|
|
83
|
+
creds = {"accountId": "acc-123", "refreshToken": "iis_rt_test"}
|
|
84
|
+
|
|
85
|
+
with (
|
|
86
|
+
patch("iis_access.auth.get_credentials", return_value=creds),
|
|
87
|
+
patch("iis_access.auth._get_keyring", return_value=None),
|
|
88
|
+
patch(
|
|
89
|
+
"iis_access.auth.get_account_url",
|
|
90
|
+
return_value="http://localhost:9060",
|
|
91
|
+
),
|
|
92
|
+
aioresponses() as mocked,
|
|
93
|
+
):
|
|
94
|
+
mocked.post(
|
|
95
|
+
"http://localhost:9060/auth/refresh",
|
|
96
|
+
status=401,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
result = await get_auth_context()
|
|
100
|
+
assert result is None
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class TestIsAuthenticated:
|
|
104
|
+
def test_returns_false_when_no_credentials(self):
|
|
105
|
+
with patch("iis_access.auth.get_credentials", return_value=None):
|
|
106
|
+
assert is_authenticated() is False
|
|
107
|
+
|
|
108
|
+
def test_returns_false_when_no_account_id(self):
|
|
109
|
+
with patch(
|
|
110
|
+
"iis_access.auth.get_credentials",
|
|
111
|
+
return_value={"refreshToken": "tok"},
|
|
112
|
+
):
|
|
113
|
+
assert is_authenticated() is False
|
|
114
|
+
|
|
115
|
+
def test_returns_true_when_credentials_exist(self):
|
|
116
|
+
with patch(
|
|
117
|
+
"iis_access.auth.get_credentials",
|
|
118
|
+
return_value={"accountId": "acc-123", "refreshToken": "tok"},
|
|
119
|
+
):
|
|
120
|
+
assert is_authenticated() is True
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Test that all functions degrade gracefully when services are unavailable."""
|
|
2
|
+
|
|
3
|
+
from unittest.mock import patch, AsyncMock
|
|
4
|
+
|
|
5
|
+
from aioresponses import aioresponses
|
|
6
|
+
|
|
7
|
+
from iis_access.auth import get_auth_context
|
|
8
|
+
from iis_access.managed_keys import get_managed_key
|
|
9
|
+
from iis_access.usage import report_usage
|
|
10
|
+
from iis_access.nexus import is_nexus_running
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class TestGracefulDegradation:
|
|
14
|
+
async def test_auth_returns_none_without_credentials(self):
|
|
15
|
+
with patch("iis_access.auth.get_credentials", return_value=None):
|
|
16
|
+
result = await get_auth_context()
|
|
17
|
+
assert result is None
|
|
18
|
+
|
|
19
|
+
async def test_auth_returns_none_on_network_timeout(self):
|
|
20
|
+
creds = {"accountId": "acc-123", "refreshToken": "iis_rt_test"}
|
|
21
|
+
|
|
22
|
+
with (
|
|
23
|
+
patch("iis_access.auth.get_credentials", return_value=creds),
|
|
24
|
+
patch("iis_access.auth._get_keyring", return_value=None),
|
|
25
|
+
patch(
|
|
26
|
+
"iis_access.auth.get_account_url",
|
|
27
|
+
return_value="http://localhost:9060",
|
|
28
|
+
),
|
|
29
|
+
aioresponses() as mocked,
|
|
30
|
+
):
|
|
31
|
+
mocked.post(
|
|
32
|
+
"http://localhost:9060/auth/refresh",
|
|
33
|
+
exception=TimeoutError("Connection timed out"),
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
result = await get_auth_context()
|
|
37
|
+
assert result is None
|
|
38
|
+
|
|
39
|
+
async def test_auth_returns_none_on_server_error(self):
|
|
40
|
+
creds = {"accountId": "acc-123", "refreshToken": "iis_rt_test"}
|
|
41
|
+
|
|
42
|
+
with (
|
|
43
|
+
patch("iis_access.auth.get_credentials", return_value=creds),
|
|
44
|
+
patch("iis_access.auth._get_keyring", return_value=None),
|
|
45
|
+
patch(
|
|
46
|
+
"iis_access.auth.get_account_url",
|
|
47
|
+
return_value="http://localhost:9060",
|
|
48
|
+
),
|
|
49
|
+
aioresponses() as mocked,
|
|
50
|
+
):
|
|
51
|
+
mocked.post(
|
|
52
|
+
"http://localhost:9060/auth/refresh",
|
|
53
|
+
status=500,
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
result = await get_auth_context()
|
|
57
|
+
assert result is None
|
|
58
|
+
|
|
59
|
+
async def test_managed_key_returns_none_without_auth(self):
|
|
60
|
+
with patch(
|
|
61
|
+
"iis_access.managed_keys.get_auth_context",
|
|
62
|
+
new_callable=AsyncMock,
|
|
63
|
+
return_value=None,
|
|
64
|
+
):
|
|
65
|
+
result = await get_managed_key("openai")
|
|
66
|
+
assert result is None
|
|
67
|
+
|
|
68
|
+
async def test_usage_silent_on_failure(self):
|
|
69
|
+
"""report_usage must never raise, even on complete failure."""
|
|
70
|
+
with patch(
|
|
71
|
+
"iis_access.usage.get_auth_context",
|
|
72
|
+
new_callable=AsyncMock,
|
|
73
|
+
return_value=None,
|
|
74
|
+
):
|
|
75
|
+
# Should complete without raising
|
|
76
|
+
await report_usage("aria", credits=1.0)
|
|
77
|
+
|
|
78
|
+
async def test_nexus_returns_false_when_not_running(self):
|
|
79
|
+
with aioresponses() as mocked:
|
|
80
|
+
mocked.get(
|
|
81
|
+
"http://localhost:9070/health",
|
|
82
|
+
exception=ConnectionError("Connection refused"),
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
result = await is_nexus_running()
|
|
86
|
+
assert result is False
|
|
87
|
+
|
|
88
|
+
async def test_nexus_returns_false_on_non_200(self):
|
|
89
|
+
with aioresponses() as mocked:
|
|
90
|
+
mocked.get(
|
|
91
|
+
"http://localhost:9070/health",
|
|
92
|
+
status=503,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
result = await is_nexus_running()
|
|
96
|
+
assert result is False
|
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
from unittest.mock import patch, AsyncMock
|
|
2
|
+
|
|
3
|
+
import pytest
|
|
4
|
+
from aioresponses import aioresponses
|
|
5
|
+
|
|
6
|
+
from iis_access._types import AuthContext
|
|
7
|
+
from iis_access.managed_keys import get_managed_key
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _make_context():
|
|
11
|
+
return AuthContext(
|
|
12
|
+
account_id="acc-123",
|
|
13
|
+
access_token="jwt-token",
|
|
14
|
+
plan="pro",
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class TestGetManagedKey:
|
|
19
|
+
async def test_returns_none_when_not_authenticated(self):
|
|
20
|
+
with patch(
|
|
21
|
+
"iis_access.managed_keys.get_auth_context",
|
|
22
|
+
new_callable=AsyncMock,
|
|
23
|
+
return_value=None,
|
|
24
|
+
):
|
|
25
|
+
result = await get_managed_key("openai")
|
|
26
|
+
assert result is None
|
|
27
|
+
|
|
28
|
+
async def test_returns_key_on_success(self):
|
|
29
|
+
with (
|
|
30
|
+
patch(
|
|
31
|
+
"iis_access.managed_keys.get_auth_context",
|
|
32
|
+
new_callable=AsyncMock,
|
|
33
|
+
return_value=_make_context(),
|
|
34
|
+
),
|
|
35
|
+
patch(
|
|
36
|
+
"iis_access.managed_keys.get_account_url",
|
|
37
|
+
return_value="http://localhost:9060",
|
|
38
|
+
),
|
|
39
|
+
aioresponses() as mocked,
|
|
40
|
+
):
|
|
41
|
+
mocked.get(
|
|
42
|
+
"http://localhost:9060/keys/managed/openai",
|
|
43
|
+
payload={"key": "sk-managed-abc"},
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
result = await get_managed_key("openai")
|
|
47
|
+
assert result == "sk-managed-abc"
|
|
48
|
+
|
|
49
|
+
async def test_returns_cached_key_on_second_call(self):
|
|
50
|
+
with (
|
|
51
|
+
patch(
|
|
52
|
+
"iis_access.managed_keys.get_auth_context",
|
|
53
|
+
new_callable=AsyncMock,
|
|
54
|
+
return_value=_make_context(),
|
|
55
|
+
),
|
|
56
|
+
patch(
|
|
57
|
+
"iis_access.managed_keys.get_account_url",
|
|
58
|
+
return_value="http://localhost:9060",
|
|
59
|
+
),
|
|
60
|
+
aioresponses() as mocked,
|
|
61
|
+
):
|
|
62
|
+
mocked.get(
|
|
63
|
+
"http://localhost:9060/keys/managed/elevenlabs",
|
|
64
|
+
payload={"key": "sk-el-abc"},
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
# First call hits API
|
|
68
|
+
result1 = await get_managed_key("elevenlabs")
|
|
69
|
+
assert result1 == "sk-el-abc"
|
|
70
|
+
|
|
71
|
+
# Second call uses cache (no mock needed)
|
|
72
|
+
result2 = await get_managed_key("elevenlabs")
|
|
73
|
+
assert result2 == "sk-el-abc"
|
|
74
|
+
|
|
75
|
+
async def test_returns_none_on_404(self):
|
|
76
|
+
with (
|
|
77
|
+
patch(
|
|
78
|
+
"iis_access.managed_keys.get_auth_context",
|
|
79
|
+
new_callable=AsyncMock,
|
|
80
|
+
return_value=_make_context(),
|
|
81
|
+
),
|
|
82
|
+
patch(
|
|
83
|
+
"iis_access.managed_keys.get_account_url",
|
|
84
|
+
return_value="http://localhost:9060",
|
|
85
|
+
),
|
|
86
|
+
aioresponses() as mocked,
|
|
87
|
+
):
|
|
88
|
+
mocked.get(
|
|
89
|
+
"http://localhost:9060/keys/managed/unknown",
|
|
90
|
+
status=404,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
result = await get_managed_key("unknown")
|
|
94
|
+
assert result is None
|
|
95
|
+
|
|
96
|
+
async def test_returns_none_on_network_error(self):
|
|
97
|
+
with (
|
|
98
|
+
patch(
|
|
99
|
+
"iis_access.managed_keys.get_auth_context",
|
|
100
|
+
new_callable=AsyncMock,
|
|
101
|
+
return_value=_make_context(),
|
|
102
|
+
),
|
|
103
|
+
patch(
|
|
104
|
+
"iis_access.managed_keys.get_account_url",
|
|
105
|
+
return_value="http://localhost:9060",
|
|
106
|
+
),
|
|
107
|
+
aioresponses() as mocked,
|
|
108
|
+
):
|
|
109
|
+
mocked.get(
|
|
110
|
+
"http://localhost:9060/keys/managed/openai",
|
|
111
|
+
exception=ConnectionError("Network unreachable"),
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
result = await get_managed_key("openai")
|
|
115
|
+
assert result is None
|
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
from unittest.mock import patch, AsyncMock
|
|
2
|
+
|
|
3
|
+
import pytest
|
|
4
|
+
|
|
5
|
+
from iis_access._types import NoProviderAvailable, ProviderResolution
|
|
6
|
+
from iis_access.provider import resolve_provider, register_free_provider
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class TestResolveProvider:
|
|
10
|
+
async def test_returns_local_key_when_env_var_set(self):
|
|
11
|
+
with (
|
|
12
|
+
patch(
|
|
13
|
+
"iis_access.provider.get_auth_context",
|
|
14
|
+
new_callable=AsyncMock,
|
|
15
|
+
return_value=None,
|
|
16
|
+
),
|
|
17
|
+
patch(
|
|
18
|
+
"iis_access.provider.is_nexus_running",
|
|
19
|
+
new_callable=AsyncMock,
|
|
20
|
+
return_value=False,
|
|
21
|
+
),
|
|
22
|
+
patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test-123"}),
|
|
23
|
+
):
|
|
24
|
+
result = await resolve_provider("llm", preferred="openai")
|
|
25
|
+
|
|
26
|
+
assert isinstance(result, ProviderResolution)
|
|
27
|
+
assert result.provider == "openai"
|
|
28
|
+
assert result.method == "local_key"
|
|
29
|
+
assert result.api_key == "sk-test-123"
|
|
30
|
+
assert result.endpoint is None
|
|
31
|
+
|
|
32
|
+
async def test_returns_free_provider_when_no_keys(self):
|
|
33
|
+
with (
|
|
34
|
+
patch(
|
|
35
|
+
"iis_access.provider.get_auth_context",
|
|
36
|
+
new_callable=AsyncMock,
|
|
37
|
+
return_value=None,
|
|
38
|
+
),
|
|
39
|
+
patch(
|
|
40
|
+
"iis_access.provider.is_nexus_running",
|
|
41
|
+
new_callable=AsyncMock,
|
|
42
|
+
return_value=False,
|
|
43
|
+
),
|
|
44
|
+
patch.dict("os.environ", {}, clear=True),
|
|
45
|
+
):
|
|
46
|
+
result = await resolve_provider("tts")
|
|
47
|
+
|
|
48
|
+
assert result.provider == "edge-tts"
|
|
49
|
+
assert result.method == "free"
|
|
50
|
+
assert result.api_key is None
|
|
51
|
+
assert result.credits_estimated == 0
|
|
52
|
+
|
|
53
|
+
async def test_raises_no_provider_when_no_options(self):
|
|
54
|
+
with (
|
|
55
|
+
patch(
|
|
56
|
+
"iis_access.provider.get_auth_context",
|
|
57
|
+
new_callable=AsyncMock,
|
|
58
|
+
return_value=None,
|
|
59
|
+
),
|
|
60
|
+
patch(
|
|
61
|
+
"iis_access.provider.is_nexus_running",
|
|
62
|
+
new_callable=AsyncMock,
|
|
63
|
+
return_value=False,
|
|
64
|
+
),
|
|
65
|
+
patch.dict("os.environ", {}, clear=True),
|
|
66
|
+
):
|
|
67
|
+
with pytest.raises(NoProviderAvailable, match="unknown-task"):
|
|
68
|
+
await resolve_provider("unknown-task", preferred="nonexistent")
|
|
69
|
+
|
|
70
|
+
async def test_prefers_nexus_for_pro_plan(self):
|
|
71
|
+
from iis_access._types import AuthContext
|
|
72
|
+
|
|
73
|
+
ctx = AuthContext(account_id="acc-1", access_token="jwt", plan="pro")
|
|
74
|
+
|
|
75
|
+
with (
|
|
76
|
+
patch(
|
|
77
|
+
"iis_access.provider.get_auth_context",
|
|
78
|
+
new_callable=AsyncMock,
|
|
79
|
+
return_value=ctx,
|
|
80
|
+
),
|
|
81
|
+
patch(
|
|
82
|
+
"iis_access.provider.is_nexus_running",
|
|
83
|
+
new_callable=AsyncMock,
|
|
84
|
+
return_value=True,
|
|
85
|
+
),
|
|
86
|
+
patch(
|
|
87
|
+
"iis_access.provider._get_nexus_endpoint",
|
|
88
|
+
return_value="http://localhost:9070",
|
|
89
|
+
),
|
|
90
|
+
):
|
|
91
|
+
result = await resolve_provider("llm", preferred="openai")
|
|
92
|
+
|
|
93
|
+
assert result.method == "nexus"
|
|
94
|
+
assert result.endpoint == "http://localhost:9070"
|
|
95
|
+
|
|
96
|
+
async def test_local_key_scan_without_preferred(self):
|
|
97
|
+
"""When no preferred provider, scans all known env vars."""
|
|
98
|
+
with (
|
|
99
|
+
patch(
|
|
100
|
+
"iis_access.provider.get_auth_context",
|
|
101
|
+
new_callable=AsyncMock,
|
|
102
|
+
return_value=None,
|
|
103
|
+
),
|
|
104
|
+
patch(
|
|
105
|
+
"iis_access.provider.is_nexus_running",
|
|
106
|
+
new_callable=AsyncMock,
|
|
107
|
+
return_value=False,
|
|
108
|
+
),
|
|
109
|
+
patch.dict(
|
|
110
|
+
"os.environ",
|
|
111
|
+
{"ANTHROPIC_API_KEY": "sk-ant-test"},
|
|
112
|
+
clear=True,
|
|
113
|
+
),
|
|
114
|
+
):
|
|
115
|
+
result = await resolve_provider("llm")
|
|
116
|
+
|
|
117
|
+
assert result.provider == "anthropic"
|
|
118
|
+
assert result.method == "local_key"
|
|
119
|
+
assert result.api_key == "sk-ant-test"
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class TestRegisterFreeProvider:
|
|
123
|
+
async def test_register_and_use_free_provider(self):
|
|
124
|
+
register_free_provider("image-gen", "local-sd", check=lambda: True)
|
|
125
|
+
|
|
126
|
+
with (
|
|
127
|
+
patch(
|
|
128
|
+
"iis_access.provider.get_auth_context",
|
|
129
|
+
new_callable=AsyncMock,
|
|
130
|
+
return_value=None,
|
|
131
|
+
),
|
|
132
|
+
patch(
|
|
133
|
+
"iis_access.provider.is_nexus_running",
|
|
134
|
+
new_callable=AsyncMock,
|
|
135
|
+
return_value=False,
|
|
136
|
+
),
|
|
137
|
+
patch.dict("os.environ", {}, clear=True),
|
|
138
|
+
):
|
|
139
|
+
result = await resolve_provider("image-gen")
|
|
140
|
+
|
|
141
|
+
assert result.provider == "local-sd"
|
|
142
|
+
assert result.method == "free"
|
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
from unittest.mock import patch, AsyncMock
|
|
2
|
+
|
|
3
|
+
from aioresponses import aioresponses
|
|
4
|
+
|
|
5
|
+
from iis_access._types import AuthContext
|
|
6
|
+
from iis_access.usage import report_usage
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class TestReportUsage:
|
|
10
|
+
async def test_does_nothing_when_not_authenticated(self):
|
|
11
|
+
with patch(
|
|
12
|
+
"iis_access.usage.get_auth_context",
|
|
13
|
+
new_callable=AsyncMock,
|
|
14
|
+
return_value=None,
|
|
15
|
+
):
|
|
16
|
+
# Should not raise
|
|
17
|
+
await report_usage("aria", credits=1.0)
|
|
18
|
+
|
|
19
|
+
async def test_sends_post_when_authenticated(self):
|
|
20
|
+
ctx = AuthContext(
|
|
21
|
+
account_id="acc-123",
|
|
22
|
+
access_token="jwt-token",
|
|
23
|
+
plan="pro",
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
with (
|
|
27
|
+
patch(
|
|
28
|
+
"iis_access.usage.get_auth_context",
|
|
29
|
+
new_callable=AsyncMock,
|
|
30
|
+
return_value=ctx,
|
|
31
|
+
),
|
|
32
|
+
patch(
|
|
33
|
+
"iis_access.usage.get_account_url",
|
|
34
|
+
return_value="http://localhost:9060",
|
|
35
|
+
),
|
|
36
|
+
aioresponses() as mocked,
|
|
37
|
+
):
|
|
38
|
+
mocked.post("http://localhost:9060/usage/record", status=200)
|
|
39
|
+
|
|
40
|
+
await report_usage("aria", credits=2.5, requests=1)
|
|
41
|
+
|
|
42
|
+
# Verify the request was made
|
|
43
|
+
requests_made = mocked.requests
|
|
44
|
+
key = ("POST", mocked.requests)
|
|
45
|
+
# aioresponses stores requests - verify at least one was made
|
|
46
|
+
assert len(requests_made) > 0
|
|
47
|
+
|
|
48
|
+
async def test_handles_network_error_gracefully(self):
|
|
49
|
+
ctx = AuthContext(
|
|
50
|
+
account_id="acc-123",
|
|
51
|
+
access_token="jwt-token",
|
|
52
|
+
plan="pro",
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
with (
|
|
56
|
+
patch(
|
|
57
|
+
"iis_access.usage.get_auth_context",
|
|
58
|
+
new_callable=AsyncMock,
|
|
59
|
+
return_value=ctx,
|
|
60
|
+
),
|
|
61
|
+
patch(
|
|
62
|
+
"iis_access.usage.get_account_url",
|
|
63
|
+
return_value="http://localhost:9060",
|
|
64
|
+
),
|
|
65
|
+
aioresponses() as mocked,
|
|
66
|
+
):
|
|
67
|
+
mocked.post(
|
|
68
|
+
"http://localhost:9060/usage/record",
|
|
69
|
+
exception=ConnectionError("Network down"),
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
# Should not raise
|
|
73
|
+
await report_usage("aria", credits=1.0)
|