citadel-predict 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- citadel_predict/__init__.py +30 -0
- citadel_predict/cli.py +179 -0
- citadel_predict/client.py +243 -0
- citadel_predict/config.py +115 -0
- citadel_predict/errors.py +109 -0
- citadel_predict/py.typed +1 -0
- citadel_predict-0.1.0.dist-info/METADATA +215 -0
- citadel_predict-0.1.0.dist-info/RECORD +10 -0
- citadel_predict-0.1.0.dist-info/WHEEL +4 -0
- citadel_predict-0.1.0.dist-info/entry_points.txt +2 -0
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""
|
|
2
|
+
citadel-predict: Developer client for AI agent pre-execution cost prediction.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from .client import CitadelClient, predict_cost
|
|
6
|
+
from .config import resolve_api_key, resolve_api_url
|
|
7
|
+
from .errors import (
|
|
8
|
+
CitadelAuthError,
|
|
9
|
+
CitadelBadRequestError,
|
|
10
|
+
CitadelError,
|
|
11
|
+
CitadelNetworkError,
|
|
12
|
+
CitadelRateLimitError,
|
|
13
|
+
CitadelServerError,
|
|
14
|
+
CitadelValidationError,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
__version__ = "0.1.0"
|
|
18
|
+
__all__ = [
|
|
19
|
+
"predict_cost",
|
|
20
|
+
"CitadelClient",
|
|
21
|
+
"resolve_api_key",
|
|
22
|
+
"resolve_api_url",
|
|
23
|
+
"CitadelError",
|
|
24
|
+
"CitadelAuthError",
|
|
25
|
+
"CitadelRateLimitError",
|
|
26
|
+
"CitadelValidationError",
|
|
27
|
+
"CitadelBadRequestError",
|
|
28
|
+
"CitadelServerError",
|
|
29
|
+
"CitadelNetworkError",
|
|
30
|
+
]
|
citadel_predict/cli.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Command-line interface for Citadel Predict.
|
|
3
|
+
|
|
4
|
+
Usage:
|
|
5
|
+
citadel-predict --task "Research 3 competitors" --tools web_search,draft_document
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import argparse
|
|
9
|
+
import json
|
|
10
|
+
import sys
|
|
11
|
+
from typing import Optional, Sequence
|
|
12
|
+
|
|
13
|
+
from .client import CitadelClient
|
|
14
|
+
from .errors import (
|
|
15
|
+
CitadelAuthError,
|
|
16
|
+
CitadelBadRequestError,
|
|
17
|
+
CitadelError,
|
|
18
|
+
CitadelNetworkError,
|
|
19
|
+
CitadelRateLimitError,
|
|
20
|
+
CitadelServerError,
|
|
21
|
+
CitadelValidationError,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
EXIT_SUCCESS = 0
|
|
25
|
+
EXIT_GENERAL_ERROR = 1
|
|
26
|
+
EXIT_VALIDATION_ERROR = 2
|
|
27
|
+
EXIT_AUTH_ERROR = 3
|
|
28
|
+
EXIT_RATE_LIMIT_ERROR = 4
|
|
29
|
+
EXIT_SERVER_ERROR = 5
|
|
30
|
+
EXIT_NETWORK_ERROR = 6
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def parse_args(args: Optional[Sequence[str]] = None) -> argparse.Namespace:
|
|
34
|
+
parser = argparse.ArgumentParser(
|
|
35
|
+
prog="citadel-predict",
|
|
36
|
+
description="Predict AI agent token budgets and costs before execution.",
|
|
37
|
+
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
|
38
|
+
)
|
|
39
|
+
parser.add_argument(
|
|
40
|
+
"--task",
|
|
41
|
+
"-t",
|
|
42
|
+
type=str,
|
|
43
|
+
help="Task description string for the agent run.",
|
|
44
|
+
)
|
|
45
|
+
parser.add_argument(
|
|
46
|
+
"positional_task",
|
|
47
|
+
nargs="?",
|
|
48
|
+
type=str,
|
|
49
|
+
help="Task description (if --task is not passed).",
|
|
50
|
+
)
|
|
51
|
+
parser.add_argument(
|
|
52
|
+
"--tools",
|
|
53
|
+
type=str,
|
|
54
|
+
default="",
|
|
55
|
+
help="Comma-separated list of tool names (e.g. 'web_search,fetch_url,draft_document').",
|
|
56
|
+
)
|
|
57
|
+
parser.add_argument(
|
|
58
|
+
"--model-id",
|
|
59
|
+
"-m",
|
|
60
|
+
type=str,
|
|
61
|
+
default="claude-sonnet",
|
|
62
|
+
help="Model calibration identifier.",
|
|
63
|
+
)
|
|
64
|
+
parser.add_argument(
|
|
65
|
+
"--api-key",
|
|
66
|
+
"-k",
|
|
67
|
+
type=str,
|
|
68
|
+
default=None,
|
|
69
|
+
help="Citadel API key (overrides CITADEL_API_KEY env var and config file).",
|
|
70
|
+
)
|
|
71
|
+
parser.add_argument(
|
|
72
|
+
"--api-url",
|
|
73
|
+
"-u",
|
|
74
|
+
type=str,
|
|
75
|
+
default=None,
|
|
76
|
+
help="Citadel API base URL (default: http://localhost:8000 or CITADEL_API_URL).",
|
|
77
|
+
)
|
|
78
|
+
parser.add_argument(
|
|
79
|
+
"--timeout",
|
|
80
|
+
type=float,
|
|
81
|
+
default=10.0,
|
|
82
|
+
help="Request timeout in seconds.",
|
|
83
|
+
)
|
|
84
|
+
parser.add_argument(
|
|
85
|
+
"--json",
|
|
86
|
+
action="store_true",
|
|
87
|
+
help="Output raw JSON instead of formatted text.",
|
|
88
|
+
)
|
|
89
|
+
return parser.parse_args(args)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def format_pretty_output(result: dict) -> str:
|
|
93
|
+
lines = [
|
|
94
|
+
"=" * 50,
|
|
95
|
+
" CITADEL PREDICT — TOKEN BUDGET ESTIMATE",
|
|
96
|
+
"=" * 50,
|
|
97
|
+
f"Model: {result.get('model_id')}",
|
|
98
|
+
f"Expected Tokens: {result.get('expected_tokens', 0):,} tokens",
|
|
99
|
+
f"Predicted Range: {result.get('low_tokens', 0):,} – {result.get('high_tokens', 0):,} tokens",
|
|
100
|
+
f"Confidence: {result.get('confidence', 'unknown').upper()}",
|
|
101
|
+
]
|
|
102
|
+
|
|
103
|
+
if result.get("out_of_distribution"):
|
|
104
|
+
lines.append(f"OOD Warning: YES ({', '.join(result.get('ood_reasons', []))})")
|
|
105
|
+
else:
|
|
106
|
+
lines.append("OOD Warning: No (In-Distribution)")
|
|
107
|
+
|
|
108
|
+
factors = result.get("driving_factors", [])
|
|
109
|
+
if factors:
|
|
110
|
+
lines.append("Driving Factors:")
|
|
111
|
+
for factor in factors:
|
|
112
|
+
lines.append(f" • {factor}")
|
|
113
|
+
|
|
114
|
+
lines.append("=" * 50)
|
|
115
|
+
return "\n".join(lines)
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def main(args: Optional[Sequence[str]] = None) -> int:
|
|
119
|
+
parsed = parse_args(args)
|
|
120
|
+
|
|
121
|
+
task = parsed.task or parsed.positional_task
|
|
122
|
+
if not task:
|
|
123
|
+
sys.stderr.write("Error: Task description is required. Use --task '...' or pass as argument.\n")
|
|
124
|
+
return EXIT_VALIDATION_ERROR
|
|
125
|
+
|
|
126
|
+
tools = [t.strip() for t in parsed.tools.split(",") if t.strip()] if parsed.tools else []
|
|
127
|
+
|
|
128
|
+
try:
|
|
129
|
+
with CitadelClient(
|
|
130
|
+
api_key=parsed.api_key,
|
|
131
|
+
base_url=parsed.api_url,
|
|
132
|
+
timeout=parsed.timeout,
|
|
133
|
+
) as client:
|
|
134
|
+
result = client.predict(
|
|
135
|
+
task_text=task,
|
|
136
|
+
tools=tools,
|
|
137
|
+
model_id=parsed.model_id,
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
if parsed.json:
|
|
141
|
+
print(json.dumps(result, indent=2))
|
|
142
|
+
else:
|
|
143
|
+
print(format_pretty_output(result))
|
|
144
|
+
return EXIT_SUCCESS
|
|
145
|
+
|
|
146
|
+
except CitadelAuthError as exc:
|
|
147
|
+
sys.stderr.write(f"Authentication Error: {exc.message}\n")
|
|
148
|
+
return EXIT_AUTH_ERROR
|
|
149
|
+
|
|
150
|
+
except CitadelRateLimitError as exc:
|
|
151
|
+
msg = f"Rate Limit Error: {exc.message}"
|
|
152
|
+
if exc.retry_after is not None:
|
|
153
|
+
msg += f" (Retry after {exc.retry_after}s)"
|
|
154
|
+
sys.stderr.write(f"{msg}\n")
|
|
155
|
+
return EXIT_RATE_LIMIT_ERROR
|
|
156
|
+
|
|
157
|
+
except (CitadelValidationError, CitadelBadRequestError) as exc:
|
|
158
|
+
sys.stderr.write(f"Validation Error: {exc.message}\n")
|
|
159
|
+
return EXIT_VALIDATION_ERROR
|
|
160
|
+
|
|
161
|
+
except CitadelServerError as exc:
|
|
162
|
+
sys.stderr.write(f"Server Error: {exc.message}\n")
|
|
163
|
+
return EXIT_SERVER_ERROR
|
|
164
|
+
|
|
165
|
+
except CitadelNetworkError as exc:
|
|
166
|
+
sys.stderr.write(f"Network Error: {exc.message}\n")
|
|
167
|
+
return EXIT_NETWORK_ERROR
|
|
168
|
+
|
|
169
|
+
except CitadelError as exc:
|
|
170
|
+
sys.stderr.write(f"Citadel Error: {exc.message}\n")
|
|
171
|
+
return EXIT_GENERAL_ERROR
|
|
172
|
+
|
|
173
|
+
except Exception as exc:
|
|
174
|
+
sys.stderr.write(f"Unexpected Error: {exc}\n")
|
|
175
|
+
return EXIT_GENERAL_ERROR
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
if __name__ == "__main__":
|
|
179
|
+
sys.exit(main())
|
|
@@ -0,0 +1,243 @@
|
|
|
1
|
+
"""
|
|
2
|
+
HTTP client for the Citadel Predict API.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any, Optional
|
|
7
|
+
|
|
8
|
+
import httpx
|
|
9
|
+
|
|
10
|
+
from .config import resolve_api_key, resolve_api_url
|
|
11
|
+
from .errors import (
|
|
12
|
+
CitadelAuthError,
|
|
13
|
+
CitadelBadRequestError,
|
|
14
|
+
CitadelError,
|
|
15
|
+
CitadelNetworkError,
|
|
16
|
+
CitadelRateLimitError,
|
|
17
|
+
CitadelServerError,
|
|
18
|
+
CitadelValidationError,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class CitadelClient:
|
|
23
|
+
"""
|
|
24
|
+
Client for interacting with the Citadel Predict API.
|
|
25
|
+
|
|
26
|
+
Handles authentication, error mapping, timeouts, and request dispatch.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self,
|
|
31
|
+
api_key: Optional[str] = None,
|
|
32
|
+
base_url: Optional[str] = None,
|
|
33
|
+
timeout: float = 10.0,
|
|
34
|
+
config_path: Optional[Path] = None,
|
|
35
|
+
http_client: Optional[httpx.Client] = None,
|
|
36
|
+
) -> None:
|
|
37
|
+
self.api_key = resolve_api_key(api_key, config_path=config_path)
|
|
38
|
+
self.base_url = resolve_api_url(base_url, config_path=config_path)
|
|
39
|
+
self.timeout = timeout
|
|
40
|
+
self._custom_client = http_client is not None
|
|
41
|
+
self._client = http_client or httpx.Client(timeout=self.timeout)
|
|
42
|
+
|
|
43
|
+
def _get_headers(self) -> dict[str, str]:
|
|
44
|
+
headers = {
|
|
45
|
+
"Content-Type": "application/json",
|
|
46
|
+
"Accept": "application/json",
|
|
47
|
+
"User-Agent": "citadel-predict-python/0.1.0",
|
|
48
|
+
}
|
|
49
|
+
if self.api_key:
|
|
50
|
+
headers["Authorization"] = f"Bearer {self.api_key}"
|
|
51
|
+
return headers
|
|
52
|
+
|
|
53
|
+
def _handle_error_response(self, response: httpx.Response) -> None:
|
|
54
|
+
status_code = response.status_code
|
|
55
|
+
try:
|
|
56
|
+
data = response.json()
|
|
57
|
+
detail = data.get("detail")
|
|
58
|
+
if isinstance(detail, list):
|
|
59
|
+
# Pydantic validation error list format
|
|
60
|
+
msg = "; ".join(
|
|
61
|
+
f"{err.get('loc', [])}: {err.get('msg', '')}" if isinstance(err, dict) else str(err)
|
|
62
|
+
for err in detail
|
|
63
|
+
)
|
|
64
|
+
elif detail:
|
|
65
|
+
msg = str(detail)
|
|
66
|
+
else:
|
|
67
|
+
msg = response.text or f"HTTP error {status_code}"
|
|
68
|
+
except Exception:
|
|
69
|
+
data = None
|
|
70
|
+
msg = response.text or f"HTTP error {status_code}"
|
|
71
|
+
|
|
72
|
+
if status_code == 401:
|
|
73
|
+
raise CitadelAuthError(
|
|
74
|
+
message=(
|
|
75
|
+
f"Authentication failed ({msg}). "
|
|
76
|
+
"Ensure your Citadel API key is set via api_key argument, "
|
|
77
|
+
"CITADEL_API_KEY environment variable, or ~/.citadel/config.toml"
|
|
78
|
+
),
|
|
79
|
+
status_code=401,
|
|
80
|
+
response_data=data,
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
if status_code == 429:
|
|
84
|
+
retry_after: Optional[int] = None
|
|
85
|
+
raw_retry = response.headers.get("Retry-After")
|
|
86
|
+
if raw_retry and raw_retry.isdigit():
|
|
87
|
+
retry_after = int(raw_retry)
|
|
88
|
+
raise CitadelRateLimitError(
|
|
89
|
+
message=f"Rate limit exceeded: {msg}",
|
|
90
|
+
status_code=429,
|
|
91
|
+
response_data=data,
|
|
92
|
+
retry_after=retry_after,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
if status_code == 422:
|
|
96
|
+
raise CitadelValidationError(
|
|
97
|
+
message=f"Validation error: {msg}",
|
|
98
|
+
status_code=422,
|
|
99
|
+
response_data=data,
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
if status_code == 400:
|
|
103
|
+
raise CitadelBadRequestError(
|
|
104
|
+
message=f"Bad request: {msg}",
|
|
105
|
+
status_code=400,
|
|
106
|
+
response_data=data,
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
if 500 <= status_code < 600:
|
|
110
|
+
raise CitadelServerError(
|
|
111
|
+
message=f"Server error ({status_code}): {msg}",
|
|
112
|
+
status_code=status_code,
|
|
113
|
+
response_data=data,
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
raise CitadelError(
|
|
117
|
+
message=f"API request failed with status {status_code}: {msg}",
|
|
118
|
+
status_code=status_code,
|
|
119
|
+
response_data=data,
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
def predict(
|
|
123
|
+
self,
|
|
124
|
+
task_text: str,
|
|
125
|
+
tools: Optional[list[str]] = None,
|
|
126
|
+
num_tools: Optional[int] = None,
|
|
127
|
+
model_id: str = "claude-sonnet",
|
|
128
|
+
) -> dict[str, Any]:
|
|
129
|
+
"""
|
|
130
|
+
Predict token budget and cost ranges for a given agent task description and tools.
|
|
131
|
+
|
|
132
|
+
Parameters:
|
|
133
|
+
task_text: Natural language task description (1-4000 chars).
|
|
134
|
+
tools: List of tool names available to the agent.
|
|
135
|
+
num_tools: Optional integer count (tools list takes precedence if provided).
|
|
136
|
+
model_id: Model calibration identifier (default: 'claude-sonnet').
|
|
137
|
+
|
|
138
|
+
Returns:
|
|
139
|
+
Dictionary matching the API response schema exactly:
|
|
140
|
+
{
|
|
141
|
+
"model_id": "claude-sonnet",
|
|
142
|
+
"low_tokens": int,
|
|
143
|
+
"expected_tokens": int,
|
|
144
|
+
"high_tokens": int,
|
|
145
|
+
"driving_factors": list[str],
|
|
146
|
+
"features": dict,
|
|
147
|
+
"confidence": str,
|
|
148
|
+
"out_of_distribution": bool,
|
|
149
|
+
"ood_reasons": list[str]
|
|
150
|
+
}
|
|
151
|
+
"""
|
|
152
|
+
tool_list = list(tools) if tools is not None else []
|
|
153
|
+
url = f"{self.base_url}/api/predict"
|
|
154
|
+
payload = {
|
|
155
|
+
"task_text": task_text,
|
|
156
|
+
"tools": tool_list,
|
|
157
|
+
"model_id": model_id,
|
|
158
|
+
}
|
|
159
|
+
|
|
160
|
+
try:
|
|
161
|
+
response = self._client.post(
|
|
162
|
+
url,
|
|
163
|
+
json=payload,
|
|
164
|
+
headers=self._get_headers(),
|
|
165
|
+
)
|
|
166
|
+
except (httpx.TimeoutException, httpx.NetworkError, httpx.RequestError) as exc:
|
|
167
|
+
raise CitadelNetworkError(
|
|
168
|
+
message=f"Network error communicating with Citadel API at {url}: {exc}",
|
|
169
|
+
cause=exc,
|
|
170
|
+
) from exc
|
|
171
|
+
|
|
172
|
+
if response.status_code != 200:
|
|
173
|
+
self._handle_error_response(response)
|
|
174
|
+
|
|
175
|
+
try:
|
|
176
|
+
return response.json()
|
|
177
|
+
except Exception as exc:
|
|
178
|
+
raise CitadelServerError(
|
|
179
|
+
message="Failed to parse JSON response from Citadel Predict API",
|
|
180
|
+
status_code=response.status_code,
|
|
181
|
+
) from exc
|
|
182
|
+
|
|
183
|
+
def health(self) -> dict[str, Any]:
|
|
184
|
+
"""Check health status and supported models on the Citadel Predict API."""
|
|
185
|
+
url = f"{self.base_url}/api/health"
|
|
186
|
+
try:
|
|
187
|
+
response = self._client.get(url, headers=self._get_headers())
|
|
188
|
+
except (httpx.TimeoutException, httpx.NetworkError, httpx.RequestError) as exc:
|
|
189
|
+
raise CitadelNetworkError(
|
|
190
|
+
message=f"Network error communicating with Citadel API at {url}: {exc}",
|
|
191
|
+
cause=exc,
|
|
192
|
+
) from exc
|
|
193
|
+
|
|
194
|
+
if response.status_code != 200:
|
|
195
|
+
self._handle_error_response(response)
|
|
196
|
+
|
|
197
|
+
return response.json()
|
|
198
|
+
|
|
199
|
+
def close(self) -> None:
|
|
200
|
+
"""Close the underlying HTTP client session."""
|
|
201
|
+
if not self._custom_client:
|
|
202
|
+
self._client.close()
|
|
203
|
+
|
|
204
|
+
def __enter__(self) -> "CitadelClient":
|
|
205
|
+
return self
|
|
206
|
+
|
|
207
|
+
def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None:
|
|
208
|
+
self.close()
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
def predict_cost(
|
|
212
|
+
task_text: str,
|
|
213
|
+
tools: Optional[list[str]] = None,
|
|
214
|
+
num_tools: Optional[int] = None,
|
|
215
|
+
model_id: str = "claude-sonnet",
|
|
216
|
+
api_key: Optional[str] = None,
|
|
217
|
+
api_url: Optional[str] = None,
|
|
218
|
+
timeout: Optional[float] = None,
|
|
219
|
+
**kwargs: Any,
|
|
220
|
+
) -> dict[str, Any]:
|
|
221
|
+
"""
|
|
222
|
+
Convenience function to predict token budget and cost ranges for an agent run.
|
|
223
|
+
|
|
224
|
+
Example:
|
|
225
|
+
>>> from citadel_predict import predict_cost
|
|
226
|
+
>>> result = predict_cost(
|
|
227
|
+
... task_text="Research competitor pricing and draft report",
|
|
228
|
+
... tools=["web_search", "draft_document"]
|
|
229
|
+
... )
|
|
230
|
+
>>> print(result["expected_tokens"])
|
|
231
|
+
"""
|
|
232
|
+
client_timeout = timeout if timeout is not None else 10.0
|
|
233
|
+
with CitadelClient(
|
|
234
|
+
api_key=api_key,
|
|
235
|
+
base_url=api_url,
|
|
236
|
+
timeout=client_timeout,
|
|
237
|
+
) as client:
|
|
238
|
+
return client.predict(
|
|
239
|
+
task_text=task_text,
|
|
240
|
+
tools=tools,
|
|
241
|
+
num_tools=num_tools,
|
|
242
|
+
model_id=model_id,
|
|
243
|
+
)
|
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Configuration resolver for Citadel Predict.
|
|
3
|
+
|
|
4
|
+
Resolution priority:
|
|
5
|
+
1. Explicit function / CLI parameter
|
|
6
|
+
2. Environment variables (`CITADEL_API_KEY`, `CITADEL_API_URL`)
|
|
7
|
+
3. Config file (`~/.citadel/config.toml`)
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import os
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
from typing import Any, Optional
|
|
13
|
+
|
|
14
|
+
try:
|
|
15
|
+
import tomllib
|
|
16
|
+
except ImportError: # pragma: no cover
|
|
17
|
+
try:
|
|
18
|
+
import tomli as tomllib # type: ignore[no-redef]
|
|
19
|
+
except ImportError:
|
|
20
|
+
tomllib = None # type: ignore[assignment]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
DEFAULT_API_URL = "http://localhost:8000"
|
|
24
|
+
DEFAULT_CONFIG_PATH = Path.home() / ".citadel" / "config.toml"
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def load_config_file(config_path: Optional[Path] = None) -> dict[str, Any]:
|
|
28
|
+
"""
|
|
29
|
+
Load configuration from ~/.citadel/config.toml if present.
|
|
30
|
+
Supports top-level keys or keys under a [default] section.
|
|
31
|
+
"""
|
|
32
|
+
path = config_path or DEFAULT_CONFIG_PATH
|
|
33
|
+
if not path.exists() or not path.is_file():
|
|
34
|
+
return {}
|
|
35
|
+
|
|
36
|
+
try:
|
|
37
|
+
if tomllib is not None:
|
|
38
|
+
with open(path, "rb") as f:
|
|
39
|
+
data = tomllib.load(f)
|
|
40
|
+
else:
|
|
41
|
+
# Fallback simple line-by-line parser for key = "value" pairs
|
|
42
|
+
data = {}
|
|
43
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
44
|
+
for line in f:
|
|
45
|
+
line = line.strip()
|
|
46
|
+
if not line or line.startswith("#") or line.startswith("["):
|
|
47
|
+
continue
|
|
48
|
+
if "=" in line:
|
|
49
|
+
k, v = line.split("=", 1)
|
|
50
|
+
k = k.strip()
|
|
51
|
+
v = v.strip().strip("\"'")
|
|
52
|
+
data[k] = v
|
|
53
|
+
|
|
54
|
+
# If data has a [default] table, merge it
|
|
55
|
+
if "default" in data and isinstance(data["default"], dict):
|
|
56
|
+
merged = dict(data["default"])
|
|
57
|
+
for k, v in data.items():
|
|
58
|
+
if k != "default":
|
|
59
|
+
merged[k] = v
|
|
60
|
+
return merged
|
|
61
|
+
return data
|
|
62
|
+
except Exception:
|
|
63
|
+
# If config file is corrupted or unreadable, ignore silently and return empty dict
|
|
64
|
+
return {}
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def resolve_api_key(
|
|
68
|
+
explicit_key: Optional[str] = None,
|
|
69
|
+
config_path: Optional[Path] = None,
|
|
70
|
+
) -> Optional[str]:
|
|
71
|
+
"""
|
|
72
|
+
Resolve API key according to priority:
|
|
73
|
+
1. Explicit argument
|
|
74
|
+
2. CITADEL_API_KEY environment variable (or legacy API_KEY)
|
|
75
|
+
3. api_key in ~/.citadel/config.toml
|
|
76
|
+
"""
|
|
77
|
+
if explicit_key:
|
|
78
|
+
return explicit_key.strip()
|
|
79
|
+
|
|
80
|
+
env_key = os.environ.get("CITADEL_API_KEY") or os.environ.get("API_KEY")
|
|
81
|
+
if env_key:
|
|
82
|
+
return env_key.strip()
|
|
83
|
+
|
|
84
|
+
file_config = load_config_file(config_path)
|
|
85
|
+
file_key = file_config.get("api_key")
|
|
86
|
+
if file_key and isinstance(file_key, str):
|
|
87
|
+
return file_key.strip()
|
|
88
|
+
|
|
89
|
+
return None
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def resolve_api_url(
|
|
93
|
+
explicit_url: Optional[str] = None,
|
|
94
|
+
config_path: Optional[Path] = None,
|
|
95
|
+
) -> str:
|
|
96
|
+
"""
|
|
97
|
+
Resolve base API URL according to priority:
|
|
98
|
+
1. Explicit argument
|
|
99
|
+
2. CITADEL_API_URL environment variable
|
|
100
|
+
3. api_url in ~/.citadel/config.toml
|
|
101
|
+
4. Default: http://localhost:8000
|
|
102
|
+
"""
|
|
103
|
+
if explicit_url:
|
|
104
|
+
return explicit_url.rstrip("/")
|
|
105
|
+
|
|
106
|
+
env_url = os.environ.get("CITADEL_API_URL")
|
|
107
|
+
if env_url:
|
|
108
|
+
return env_url.rstrip("/")
|
|
109
|
+
|
|
110
|
+
file_config = load_config_file(config_path)
|
|
111
|
+
file_url = file_config.get("api_url")
|
|
112
|
+
if file_url and isinstance(file_url, str):
|
|
113
|
+
return file_url.rstrip("/")
|
|
114
|
+
|
|
115
|
+
return DEFAULT_API_URL
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Exception hierarchy for citadel-predict.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from typing import Any, Optional
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class CitadelError(Exception):
|
|
9
|
+
"""Base exception for all Citadel Predict client errors."""
|
|
10
|
+
|
|
11
|
+
def __init__(
|
|
12
|
+
self,
|
|
13
|
+
message: str,
|
|
14
|
+
status_code: Optional[int] = None,
|
|
15
|
+
response_data: Optional[Any] = None,
|
|
16
|
+
) -> None:
|
|
17
|
+
super().__init__(message)
|
|
18
|
+
self.message = message
|
|
19
|
+
self.status_code = status_code
|
|
20
|
+
self.response_data = response_data
|
|
21
|
+
|
|
22
|
+
def __str__(self) -> str:
|
|
23
|
+
if self.status_code:
|
|
24
|
+
return f"[{self.status_code}] {self.message}"
|
|
25
|
+
return self.message
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class CitadelAuthError(CitadelError):
|
|
29
|
+
"""Raised when authentication fails (HTTP 401) or API key is missing/invalid."""
|
|
30
|
+
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
message: str = (
|
|
34
|
+
"Authentication failed: Missing or invalid Citadel API key. "
|
|
35
|
+
"Set CITADEL_API_KEY environment variable, pass api_key='...', "
|
|
36
|
+
"or configure api_key in ~/.citadel/config.toml"
|
|
37
|
+
),
|
|
38
|
+
status_code: int = 401,
|
|
39
|
+
response_data: Optional[Any] = None,
|
|
40
|
+
) -> None:
|
|
41
|
+
super().__init__(message, status_code=status_code, response_data=response_data)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class CitadelRateLimitError(CitadelError):
|
|
45
|
+
"""Raised when API rate limits are exceeded (HTTP 429)."""
|
|
46
|
+
|
|
47
|
+
def __init__(
|
|
48
|
+
self,
|
|
49
|
+
message: str = "Rate limit exceeded. Please back off and retry later.",
|
|
50
|
+
status_code: int = 429,
|
|
51
|
+
response_data: Optional[Any] = None,
|
|
52
|
+
retry_after: Optional[int] = None,
|
|
53
|
+
) -> None:
|
|
54
|
+
super().__init__(message, status_code=status_code, response_data=response_data)
|
|
55
|
+
self.retry_after = retry_after
|
|
56
|
+
|
|
57
|
+
def __str__(self) -> str:
|
|
58
|
+
base = f"[{self.status_code}] {self.message}"
|
|
59
|
+
if self.retry_after is not None:
|
|
60
|
+
base += f" (Retry-After: {self.retry_after}s)"
|
|
61
|
+
return base
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class CitadelValidationError(CitadelError):
|
|
65
|
+
"""Raised when the server rejects invalid request data (HTTP 422)."""
|
|
66
|
+
|
|
67
|
+
def __init__(
|
|
68
|
+
self,
|
|
69
|
+
message: str = "Request validation failed. Check task_text and tools formatting.",
|
|
70
|
+
status_code: int = 422,
|
|
71
|
+
response_data: Optional[Any] = None,
|
|
72
|
+
) -> None:
|
|
73
|
+
super().__init__(message, status_code=status_code, response_data=response_data)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class CitadelBadRequestError(CitadelError):
|
|
77
|
+
"""Raised when the request is bad or model_id is unsupported (HTTP 400)."""
|
|
78
|
+
|
|
79
|
+
def __init__(
|
|
80
|
+
self,
|
|
81
|
+
message: str = "Bad request. Unsupported model or invalid payload.",
|
|
82
|
+
status_code: int = 400,
|
|
83
|
+
response_data: Optional[Any] = None,
|
|
84
|
+
) -> None:
|
|
85
|
+
super().__init__(message, status_code=status_code, response_data=response_data)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class CitadelServerError(CitadelError):
|
|
89
|
+
"""Raised when the Citadel API server returns a 5xx error."""
|
|
90
|
+
|
|
91
|
+
def __init__(
|
|
92
|
+
self,
|
|
93
|
+
message: str = "Citadel Predict server error. Please try again later.",
|
|
94
|
+
status_code: int = 500,
|
|
95
|
+
response_data: Optional[Any] = None,
|
|
96
|
+
) -> None:
|
|
97
|
+
super().__init__(message, status_code=status_code, response_data=response_data)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
class CitadelNetworkError(CitadelError):
|
|
101
|
+
"""Raised when a network connection fails or requests time out."""
|
|
102
|
+
|
|
103
|
+
def __init__(
|
|
104
|
+
self,
|
|
105
|
+
message: str = "Network error: Failed to connect to Citadel Predict API.",
|
|
106
|
+
cause: Optional[Exception] = None,
|
|
107
|
+
) -> None:
|
|
108
|
+
super().__init__(message)
|
|
109
|
+
self.cause = cause
|
citadel_predict/py.typed
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# Marker file for PEP 561.
|
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: citadel-predict
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Pre-execution LLM token budget and cost prediction client for AI agents
|
|
5
|
+
Author: Citadel Predict Team
|
|
6
|
+
License: MIT
|
|
7
|
+
Keywords: agent,ai,budget,cost-estimation,governance,guardrails,llm,tokens
|
|
8
|
+
Classifier: Development Status :: 4 - Beta
|
|
9
|
+
Classifier: Intended Audience :: Developers
|
|
10
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
11
|
+
Classifier: Programming Language :: Python :: 3
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
16
|
+
Classifier: Topic :: Software Development :: Libraries :: Python Modules
|
|
17
|
+
Requires-Python: >=3.9
|
|
18
|
+
Requires-Dist: httpx>=0.24.0
|
|
19
|
+
Provides-Extra: dev
|
|
20
|
+
Requires-Dist: mypy>=1.10.0; extra == 'dev'
|
|
21
|
+
Requires-Dist: pytest-cov>=4.1.0; extra == 'dev'
|
|
22
|
+
Requires-Dist: pytest>=8.0.0; extra == 'dev'
|
|
23
|
+
Requires-Dist: ruff>=0.4.0; extra == 'dev'
|
|
24
|
+
Description-Content-Type: text/markdown
|
|
25
|
+
|
|
26
|
+
# citadel-predict
|
|
27
|
+
|
|
28
|
+
[](https://pypi.org/project/citadel-predict/)
|
|
29
|
+
[](https://pypi.org/project/citadel-predict/)
|
|
30
|
+
[](https://opensource.org/licenses/MIT)
|
|
31
|
+
|
|
32
|
+
**Pre-execution token budget and cost predictor client for AI agents.**
|
|
33
|
+
|
|
34
|
+
`citadel-predict` is a lightweight, pure HTTP Python client and CLI for the Citadel Predict API. It allows developers, CI pipelines, and autonomous agent loops to estimate LLM token consumption and cost ranges *before* initiating expensive agent runs.
|
|
35
|
+
|
|
36
|
+
---
|
|
37
|
+
|
|
38
|
+
## Installation
|
|
39
|
+
|
|
40
|
+
```bash
|
|
41
|
+
pip install citadel-predict
|
|
42
|
+
```
|
|
43
|
+
|
|
44
|
+
---
|
|
45
|
+
|
|
46
|
+
## Quickstart (3 Lines of Code)
|
|
47
|
+
|
|
48
|
+
```python
|
|
49
|
+
from citadel_predict import predict_cost
|
|
50
|
+
|
|
51
|
+
result = predict_cost(
|
|
52
|
+
task_text="Research competitor pricing across 3 sources and draft report",
|
|
53
|
+
tools=["web_search", "draft_document"]
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
print(f"Expected: {result['expected_tokens']:,} tokens (Range: {result['low_tokens']:,} – {result['high_tokens']:,})")
|
|
57
|
+
```
|
|
58
|
+
|
|
59
|
+
Output:
|
|
60
|
+
```text
|
|
61
|
+
Expected: 3,200 tokens (Range: 1,500 – 5,800)
|
|
62
|
+
```
|
|
63
|
+
|
|
64
|
+
---
|
|
65
|
+
|
|
66
|
+
## CLI Usage
|
|
67
|
+
|
|
68
|
+
`citadel-predict` includes a full-featured CLI for terminal workflows and CI/CD cost checks:
|
|
69
|
+
|
|
70
|
+
```bash
|
|
71
|
+
# Pretty terminal card output
|
|
72
|
+
citadel-predict --task "Audit repository and write migration guide" --tools list_files,read_document,draft_document
|
|
73
|
+
|
|
74
|
+
# Scripting / CI mode (JSON output)
|
|
75
|
+
citadel-predict --task "Calculate statistical metrics" --tools calculator --json
|
|
76
|
+
|
|
77
|
+
# Override API key or URL
|
|
78
|
+
citadel-predict --task "..." --api-key "cp_live_12345" --api-url "https://api.citadel.dev"
|
|
79
|
+
```
|
|
80
|
+
|
|
81
|
+
### CLI Exit Codes
|
|
82
|
+
- `0`: Success
|
|
83
|
+
- `2`: Validation Error / Bad Request (HTTP 400 / 422 or missing task)
|
|
84
|
+
- `3`: Authentication Failure (HTTP 401)
|
|
85
|
+
- `4`: Rate Limit Exceeded (HTTP 429)
|
|
86
|
+
- `5`: Server Error (HTTP 5xx)
|
|
87
|
+
- `6`: Network / Timeout Error
|
|
88
|
+
|
|
89
|
+
---
|
|
90
|
+
|
|
91
|
+
## Real Agent Integration: Pre-Execution Guardrails
|
|
92
|
+
|
|
93
|
+
Existing agent governance tools (e.g., Portkey, Langfuse, LiteLLM) are **reactive**—they record costs during or after an execution. `citadel-predict` is **predictive**—enabling pre-flight budget checks and dynamic routing before running reasoning loops.
|
|
94
|
+
|
|
95
|
+
### LangGraph / CrewAI Pre-Flight Cost Guardrail Example
|
|
96
|
+
|
|
97
|
+
```python
|
|
98
|
+
from typing import TypedDict, List
|
|
99
|
+
from citadel_predict import predict_cost, CitadelError
|
|
100
|
+
|
|
101
|
+
class AgentState(TypedDict):
|
|
102
|
+
task: str
|
|
103
|
+
tools: List[str]
|
|
104
|
+
budget_tokens: int
|
|
105
|
+
approved: bool
|
|
106
|
+
|
|
107
|
+
def pre_flight_budget_guardrail(state: AgentState) -> AgentState:
|
|
108
|
+
"""
|
|
109
|
+
Evaluates token budget before dispatching tools or multi-agent loops.
|
|
110
|
+
"""
|
|
111
|
+
try:
|
|
112
|
+
prediction = predict_cost(
|
|
113
|
+
task_text=state["task"],
|
|
114
|
+
tools=state["tools"],
|
|
115
|
+
model_id="claude-sonnet"
|
|
116
|
+
)
|
|
117
|
+
except CitadelError as e:
|
|
118
|
+
print(f"Cost prediction unavailable: {e}. Falling back to default budget.")
|
|
119
|
+
return state
|
|
120
|
+
|
|
121
|
+
expected = prediction["expected_tokens"]
|
|
122
|
+
high = prediction["high_tokens"]
|
|
123
|
+
is_ood = prediction["out_of_distribution"]
|
|
124
|
+
|
|
125
|
+
print(f"Pre-flight estimate: ~{expected:,} tokens (Upper bound: {high:,})")
|
|
126
|
+
if is_ood:
|
|
127
|
+
print(f"Warning: Out-of-Distribution task ({prediction['ood_reasons']})")
|
|
128
|
+
|
|
129
|
+
# Guardrail Policy: Escalate if upper bound exceeds budget
|
|
130
|
+
if high > state["budget_tokens"]:
|
|
131
|
+
print(f"[BLOCKED] High-estimate ({high:,}) exceeds budget ({state['budget_tokens']:,})")
|
|
132
|
+
# In a real agent: switch to smaller model, ask human for approval, or prune tool access
|
|
133
|
+
state["approved"] = False
|
|
134
|
+
else:
|
|
135
|
+
state["approved"] = True
|
|
136
|
+
|
|
137
|
+
return state
|
|
138
|
+
|
|
139
|
+
# Example usage in workflow
|
|
140
|
+
initial_state: AgentState = {
|
|
141
|
+
"task": "Perform exhaustive market research across 20 industry filings",
|
|
142
|
+
"tools": ["web_search", "fetch_url", "draft_document"],
|
|
143
|
+
"budget_tokens": 10000,
|
|
144
|
+
"approved": False
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
state = pre_flight_budget_guardrail(initial_state)
|
|
148
|
+
if not state["approved"]:
|
|
149
|
+
print("Action required: Human-in-the-loop approval or task reformulation needed.")
|
|
150
|
+
```
|
|
151
|
+
|
|
152
|
+
---
|
|
153
|
+
|
|
154
|
+
## Authentication & Configuration
|
|
155
|
+
|
|
156
|
+
The client resolves your API key and base URL according to the following priority:
|
|
157
|
+
|
|
158
|
+
1. **Explicit argument**: `predict_cost(..., api_key="...", api_url="...")` or CLI `--api-key` / `--api-url`
|
|
159
|
+
2. **Environment variables**: `CITADEL_API_KEY` and `CITADEL_API_URL`
|
|
160
|
+
3. **Configuration file**: `~/.citadel/config.toml`
|
|
161
|
+
|
|
162
|
+
### Example `~/.citadel/config.toml`
|
|
163
|
+
```toml
|
|
164
|
+
api_key = "cp_live_your_api_key_here"
|
|
165
|
+
api_url = "https://api.citadel.dev"
|
|
166
|
+
```
|
|
167
|
+
|
|
168
|
+
---
|
|
169
|
+
|
|
170
|
+
## Error Handling
|
|
171
|
+
|
|
172
|
+
`citadel-predict` surfaces typed, catchable exceptions:
|
|
173
|
+
|
|
174
|
+
```python
|
|
175
|
+
from citadel_predict import (
|
|
176
|
+
predict_cost,
|
|
177
|
+
CitadelAuthError,
|
|
178
|
+
CitadelRateLimitError,
|
|
179
|
+
CitadelValidationError,
|
|
180
|
+
CitadelServerError,
|
|
181
|
+
CitadelNetworkError,
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
try:
|
|
185
|
+
result = predict_cost("Analyze dataset", tools=["calculator"])
|
|
186
|
+
except CitadelAuthError:
|
|
187
|
+
# 401: Missing or invalid API key
|
|
188
|
+
...
|
|
189
|
+
except CitadelRateLimitError as e:
|
|
190
|
+
# 429: Rate limited; check e.retry_after
|
|
191
|
+
print(f"Retry after {e.retry_after} seconds")
|
|
192
|
+
except CitadelValidationError as e:
|
|
193
|
+
# 422: Input validation bounds exceeded (e.g. task > 4000 chars)
|
|
194
|
+
...
|
|
195
|
+
except CitadelNetworkError as e:
|
|
196
|
+
# Timeout or connection failure
|
|
197
|
+
...
|
|
198
|
+
```
|
|
199
|
+
|
|
200
|
+
---
|
|
201
|
+
|
|
202
|
+
## Honest Limitations
|
|
203
|
+
|
|
204
|
+
`citadel-predict` is a thin client wrapping the hosted calibration model. It directly inherits the current system characteristics:
|
|
205
|
+
|
|
206
|
+
1. **Single-Model Calibration**: Calibration is currently tuned specifically for **Claude Sonnet** (`claude-sonnet`). Future releases will introduce multi-model support via `model_id`.
|
|
207
|
+
2. **Calibration Dataset Scale**: Calibrated on $N=20$ diverse task archetypes across 80 benchmarked runs.
|
|
208
|
+
3. **Synthetic Tool Sizing**: Ground-truth data was collected using deterministic mock tool outputs with representative context expansion. Real-world tools with unbounded payload returns (e.g., massive scraped DOMs) may exhibit higher variance.
|
|
209
|
+
4. **Pre-execution Estimation**: Token predictions represent calibrated statistical ranges $[low, expected, high]$, not runtime guarantees against infinite loops or divergent agent reasoning.
|
|
210
|
+
|
|
211
|
+
---
|
|
212
|
+
|
|
213
|
+
## License
|
|
214
|
+
|
|
215
|
+
MIT
|
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
citadel_predict/__init__.py,sha256=wWIouGF5_qoLlC1abmINIPMPs8CfxqXlUwLw2yf_OUM,694
|
|
2
|
+
citadel_predict/cli.py,sha256=axeoyd-G_i4COPVPce7f6z8Utx-kUzapNJPmZSabZlA,5166
|
|
3
|
+
citadel_predict/client.py,sha256=JprRpeRpJI6gjqA3sk18F_CxP8tFc4nX7RKyn6eXM9U,8038
|
|
4
|
+
citadel_predict/config.py,sha256=P7eWXpJpjwhMSdmOlFThEtj3q7bxeE5NZ4jebgUpzt8,3431
|
|
5
|
+
citadel_predict/errors.py,sha256=0k7sY5doKBw99_BB8sYgMGM6x_7I3Zr4aXHPfT4zxfU,3495
|
|
6
|
+
citadel_predict/py.typed,sha256=bWew9mHgMy8LqMu7RuqQXFXLBxh2CRx0dUbSx-3wE48,27
|
|
7
|
+
citadel_predict-0.1.0.dist-info/METADATA,sha256=yvoo-xlKDTBRnlynovI3kWxg12-hwrrjgadHzvDX5VY,7232
|
|
8
|
+
citadel_predict-0.1.0.dist-info/WHEEL,sha256=zOwg4jB6zX2kU910N-cMawjivD6tO8NEWvE12je1bVk,87
|
|
9
|
+
citadel_predict-0.1.0.dist-info/entry_points.txt,sha256=wpobzRsGxw0Tv6f0mWFlpH7OEXpOM7D2zyed6L60GXI,61
|
|
10
|
+
citadel_predict-0.1.0.dist-info/RECORD,,
|