smartflow-sdk 0.2.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.
- smartflow/__init__.py +129 -0
- smartflow/client.py +1160 -0
- smartflow/exceptions.py +120 -0
- smartflow/sync.py +386 -0
- smartflow/types.py +526 -0
- smartflow_sdk-0.2.0.dist-info/METADATA +345 -0
- smartflow_sdk-0.2.0.dist-info/RECORD +10 -0
- smartflow_sdk-0.2.0.dist-info/WHEEL +5 -0
- smartflow_sdk-0.2.0.dist-info/licenses/LICENSE +22 -0
- smartflow_sdk-0.2.0.dist-info/top_level.txt +1 -0
smartflow/client.py
ADDED
|
@@ -0,0 +1,1160 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Smartflow SDK Client.
|
|
3
|
+
|
|
4
|
+
The main async client for connecting to a deployed Smartflow instance.
|
|
5
|
+
|
|
6
|
+
Example:
|
|
7
|
+
>>> from smartflow import SmartflowClient
|
|
8
|
+
>>>
|
|
9
|
+
>>> async with SmartflowClient("http://smartflow:7775") as sf:
|
|
10
|
+
... response = await sf.chat("What is machine learning?")
|
|
11
|
+
... print(response)
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from typing import Optional, Dict, List, Any, AsyncIterator
|
|
15
|
+
import httpx
|
|
16
|
+
import json
|
|
17
|
+
from datetime import datetime
|
|
18
|
+
|
|
19
|
+
from .types import (
|
|
20
|
+
AIRequest,
|
|
21
|
+
AIResponse,
|
|
22
|
+
ChatMessage,
|
|
23
|
+
ComplianceResult,
|
|
24
|
+
CacheStats,
|
|
25
|
+
ProviderHealth,
|
|
26
|
+
VASLog,
|
|
27
|
+
PolicyResult,
|
|
28
|
+
SystemHealth,
|
|
29
|
+
IntelligentScanResult,
|
|
30
|
+
LearningStatus,
|
|
31
|
+
LearningSummary,
|
|
32
|
+
MLStats,
|
|
33
|
+
OrgBaseline,
|
|
34
|
+
PersistenceStats,
|
|
35
|
+
AgentConfig,
|
|
36
|
+
WorkflowStep,
|
|
37
|
+
WorkflowResult,
|
|
38
|
+
)
|
|
39
|
+
from .exceptions import (
|
|
40
|
+
SmartflowError,
|
|
41
|
+
ConnectionError as SmartflowConnectionError,
|
|
42
|
+
raise_for_status,
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class SmartflowClient:
|
|
47
|
+
"""
|
|
48
|
+
Async client for connecting to a deployed Smartflow instance.
|
|
49
|
+
|
|
50
|
+
Smartflow provides:
|
|
51
|
+
- Intelligent routing across AI providers (OpenAI, Anthropic, Gemini, etc.)
|
|
52
|
+
- 3-layer semantic caching (60-80% cost savings)
|
|
53
|
+
- Real-time compliance scanning
|
|
54
|
+
- Full audit logging (VAS logs)
|
|
55
|
+
- Provider failover and load balancing
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
base_url: URL of Smartflow proxy (default port 7775)
|
|
59
|
+
api_key: Optional API key for authentication
|
|
60
|
+
timeout: Request timeout in seconds (default 30)
|
|
61
|
+
management_port: Port for management API (default 7778)
|
|
62
|
+
compliance_port: Port for compliance API (default 7777)
|
|
63
|
+
bridge_port: Port for hybrid bridge API (default 3500)
|
|
64
|
+
|
|
65
|
+
Example:
|
|
66
|
+
>>> async with SmartflowClient("http://192.81.214.94:7775") as sf:
|
|
67
|
+
... # Simple chat
|
|
68
|
+
... response = await sf.chat("What is AI?")
|
|
69
|
+
... print(response)
|
|
70
|
+
...
|
|
71
|
+
... # Check cache stats
|
|
72
|
+
... stats = await sf.get_cache_stats()
|
|
73
|
+
... print(f"Hit rate: {stats.hit_rate:.1%}")
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
def __init__(
|
|
77
|
+
self,
|
|
78
|
+
base_url: str,
|
|
79
|
+
api_key: Optional[str] = None,
|
|
80
|
+
timeout: float = 30.0,
|
|
81
|
+
management_port: int = 7778,
|
|
82
|
+
compliance_port: int = 7777,
|
|
83
|
+
bridge_port: int = 3500,
|
|
84
|
+
):
|
|
85
|
+
self.base_url = base_url.rstrip('/')
|
|
86
|
+
self.api_key = api_key
|
|
87
|
+
self.timeout = timeout
|
|
88
|
+
|
|
89
|
+
# Extract host from base_url for other ports
|
|
90
|
+
# e.g., "http://192.81.214.94:7775" -> "http://192.81.214.94"
|
|
91
|
+
parts = self.base_url.rsplit(':', 1)
|
|
92
|
+
if len(parts) == 2 and parts[1].isdigit():
|
|
93
|
+
self._host = parts[0]
|
|
94
|
+
else:
|
|
95
|
+
self._host = self.base_url
|
|
96
|
+
|
|
97
|
+
self.management_url = f"{self._host}:{management_port}"
|
|
98
|
+
self.compliance_url = f"{self._host}:{compliance_port}"
|
|
99
|
+
self.bridge_url = f"{self._host}:{bridge_port}"
|
|
100
|
+
|
|
101
|
+
self._client: Optional[httpx.AsyncClient] = None
|
|
102
|
+
|
|
103
|
+
async def _ensure_client(self):
|
|
104
|
+
"""Ensure HTTP client is initialized."""
|
|
105
|
+
if self._client is None:
|
|
106
|
+
self._client = httpx.AsyncClient(timeout=self.timeout)
|
|
107
|
+
|
|
108
|
+
def _headers(self, extra: Dict[str, str] = None) -> Dict[str, str]:
|
|
109
|
+
"""Build request headers."""
|
|
110
|
+
headers = {"Content-Type": "application/json"}
|
|
111
|
+
if self.api_key:
|
|
112
|
+
headers["Authorization"] = f"Bearer {self.api_key}"
|
|
113
|
+
if extra:
|
|
114
|
+
headers.update(extra)
|
|
115
|
+
return headers
|
|
116
|
+
|
|
117
|
+
async def _get(
|
|
118
|
+
self,
|
|
119
|
+
url: str,
|
|
120
|
+
params: Dict[str, Any] = None,
|
|
121
|
+
headers: Dict[str, str] = None,
|
|
122
|
+
) -> Dict[str, Any]:
|
|
123
|
+
"""Make GET request."""
|
|
124
|
+
await self._ensure_client()
|
|
125
|
+
try:
|
|
126
|
+
response = await self._client.get(
|
|
127
|
+
url,
|
|
128
|
+
params=params,
|
|
129
|
+
headers=self._headers(headers)
|
|
130
|
+
)
|
|
131
|
+
data = response.json()
|
|
132
|
+
raise_for_status(data, response.status_code)
|
|
133
|
+
return data
|
|
134
|
+
except httpx.ConnectError as e:
|
|
135
|
+
raise SmartflowConnectionError(f"Failed to connect to {url}: {e}")
|
|
136
|
+
|
|
137
|
+
async def _post(
|
|
138
|
+
self,
|
|
139
|
+
url: str,
|
|
140
|
+
payload: Dict[str, Any],
|
|
141
|
+
headers: Dict[str, str] = None,
|
|
142
|
+
) -> Dict[str, Any]:
|
|
143
|
+
"""Make POST request."""
|
|
144
|
+
await self._ensure_client()
|
|
145
|
+
try:
|
|
146
|
+
response = await self._client.post(
|
|
147
|
+
url,
|
|
148
|
+
json=payload,
|
|
149
|
+
headers=self._headers(headers)
|
|
150
|
+
)
|
|
151
|
+
data = response.json()
|
|
152
|
+
raise_for_status(data, response.status_code)
|
|
153
|
+
return data
|
|
154
|
+
except httpx.ConnectError as e:
|
|
155
|
+
raise SmartflowConnectionError(f"Failed to connect to {url}: {e}")
|
|
156
|
+
|
|
157
|
+
# =========================================================================
|
|
158
|
+
# CORE AI METHODS
|
|
159
|
+
# =========================================================================
|
|
160
|
+
|
|
161
|
+
async def chat(
|
|
162
|
+
self,
|
|
163
|
+
message: str,
|
|
164
|
+
model: str = "gpt-4o",
|
|
165
|
+
system_prompt: Optional[str] = None,
|
|
166
|
+
temperature: float = 0.7,
|
|
167
|
+
max_tokens: Optional[int] = None,
|
|
168
|
+
**kwargs,
|
|
169
|
+
) -> str:
|
|
170
|
+
"""
|
|
171
|
+
Send a chat message through Smartflow.
|
|
172
|
+
|
|
173
|
+
Smartflow automatically handles:
|
|
174
|
+
- Provider routing (best model selection)
|
|
175
|
+
- Semantic caching (60-80% cost savings)
|
|
176
|
+
- Compliance checking
|
|
177
|
+
- Audit logging
|
|
178
|
+
|
|
179
|
+
Args:
|
|
180
|
+
message: User message
|
|
181
|
+
model: Model to use (default "gpt-4o")
|
|
182
|
+
system_prompt: Optional system prompt
|
|
183
|
+
temperature: Sampling temperature (default 0.7)
|
|
184
|
+
max_tokens: Maximum tokens to generate
|
|
185
|
+
**kwargs: Additional parameters passed to provider
|
|
186
|
+
|
|
187
|
+
Returns:
|
|
188
|
+
AI response text
|
|
189
|
+
|
|
190
|
+
Example:
|
|
191
|
+
>>> response = await sf.chat("Explain quantum computing")
|
|
192
|
+
>>> print(response)
|
|
193
|
+
"""
|
|
194
|
+
messages = []
|
|
195
|
+
if system_prompt:
|
|
196
|
+
messages.append({"role": "system", "content": system_prompt})
|
|
197
|
+
messages.append({"role": "user", "content": message})
|
|
198
|
+
|
|
199
|
+
response = await self.chat_completions(
|
|
200
|
+
messages=messages,
|
|
201
|
+
model=model,
|
|
202
|
+
temperature=temperature,
|
|
203
|
+
max_tokens=max_tokens,
|
|
204
|
+
**kwargs,
|
|
205
|
+
)
|
|
206
|
+
return response.content
|
|
207
|
+
|
|
208
|
+
async def chat_completions(
|
|
209
|
+
self,
|
|
210
|
+
messages: List[Dict[str, str]],
|
|
211
|
+
model: str = "gpt-4o",
|
|
212
|
+
temperature: float = 0.7,
|
|
213
|
+
max_tokens: Optional[int] = None,
|
|
214
|
+
stream: bool = False,
|
|
215
|
+
**kwargs,
|
|
216
|
+
) -> AIResponse:
|
|
217
|
+
"""
|
|
218
|
+
OpenAI-compatible chat completions endpoint.
|
|
219
|
+
|
|
220
|
+
Fully compatible with existing OpenAI code - just change the base URL!
|
|
221
|
+
|
|
222
|
+
Args:
|
|
223
|
+
messages: List of message dicts with "role" and "content"
|
|
224
|
+
model: Model to use
|
|
225
|
+
temperature: Sampling temperature
|
|
226
|
+
max_tokens: Maximum tokens
|
|
227
|
+
stream: Enable streaming (returns async iterator if True)
|
|
228
|
+
**kwargs: Additional parameters
|
|
229
|
+
|
|
230
|
+
Returns:
|
|
231
|
+
AIResponse object with choices, usage, etc.
|
|
232
|
+
|
|
233
|
+
Example:
|
|
234
|
+
>>> response = await sf.chat_completions(
|
|
235
|
+
... messages=[{"role": "user", "content": "Hello!"}],
|
|
236
|
+
... model="gpt-4o"
|
|
237
|
+
... )
|
|
238
|
+
>>> print(response.content)
|
|
239
|
+
"""
|
|
240
|
+
payload = {
|
|
241
|
+
"model": model,
|
|
242
|
+
"messages": messages,
|
|
243
|
+
"temperature": temperature,
|
|
244
|
+
"stream": stream,
|
|
245
|
+
**kwargs,
|
|
246
|
+
}
|
|
247
|
+
if max_tokens:
|
|
248
|
+
payload["max_tokens"] = max_tokens
|
|
249
|
+
|
|
250
|
+
url = f"{self.base_url}/v1/chat/completions"
|
|
251
|
+
data = await self._post(url, payload)
|
|
252
|
+
return AIResponse.from_dict(data)
|
|
253
|
+
|
|
254
|
+
async def embeddings(
|
|
255
|
+
self,
|
|
256
|
+
input: str | List[str],
|
|
257
|
+
model: str = "text-embedding-3-small",
|
|
258
|
+
) -> Dict[str, Any]:
|
|
259
|
+
"""
|
|
260
|
+
Generate embeddings for text.
|
|
261
|
+
|
|
262
|
+
Args:
|
|
263
|
+
input: Text or list of texts to embed
|
|
264
|
+
model: Embedding model to use
|
|
265
|
+
|
|
266
|
+
Returns:
|
|
267
|
+
Dict with embeddings data
|
|
268
|
+
"""
|
|
269
|
+
payload = {
|
|
270
|
+
"model": model,
|
|
271
|
+
"input": input if isinstance(input, list) else [input],
|
|
272
|
+
}
|
|
273
|
+
url = f"{self.base_url}/v1/embeddings"
|
|
274
|
+
return await self._post(url, payload)
|
|
275
|
+
|
|
276
|
+
async def list_models(self) -> List[Dict[str, Any]]:
|
|
277
|
+
"""
|
|
278
|
+
List available models.
|
|
279
|
+
|
|
280
|
+
Returns:
|
|
281
|
+
List of model info dicts
|
|
282
|
+
"""
|
|
283
|
+
url = f"{self.base_url}/v1/models"
|
|
284
|
+
data = await self._get(url)
|
|
285
|
+
return data.get("data", [])
|
|
286
|
+
|
|
287
|
+
# =========================================================================
|
|
288
|
+
# ANTHROPIC METHODS
|
|
289
|
+
# =========================================================================
|
|
290
|
+
|
|
291
|
+
async def claude_message(
|
|
292
|
+
self,
|
|
293
|
+
message: str,
|
|
294
|
+
model: str = "claude-3-5-sonnet-20241022",
|
|
295
|
+
max_tokens: int = 1024,
|
|
296
|
+
system: Optional[str] = None,
|
|
297
|
+
anthropic_key: Optional[str] = None,
|
|
298
|
+
) -> str:
|
|
299
|
+
"""
|
|
300
|
+
Send a message to Claude via Anthropic API.
|
|
301
|
+
|
|
302
|
+
Args:
|
|
303
|
+
message: User message
|
|
304
|
+
model: Claude model to use
|
|
305
|
+
max_tokens: Maximum tokens
|
|
306
|
+
system: System prompt
|
|
307
|
+
anthropic_key: Anthropic API key (uses stored key if not provided)
|
|
308
|
+
|
|
309
|
+
Returns:
|
|
310
|
+
Claude's response text
|
|
311
|
+
"""
|
|
312
|
+
payload = {
|
|
313
|
+
"model": model,
|
|
314
|
+
"max_tokens": max_tokens,
|
|
315
|
+
"messages": [{"role": "user", "content": message}],
|
|
316
|
+
}
|
|
317
|
+
if system:
|
|
318
|
+
payload["system"] = system
|
|
319
|
+
|
|
320
|
+
headers = {
|
|
321
|
+
"anthropic-version": "2023-06-01",
|
|
322
|
+
}
|
|
323
|
+
if anthropic_key:
|
|
324
|
+
headers["x-api-key"] = anthropic_key
|
|
325
|
+
|
|
326
|
+
url = f"{self.base_url}/v1/messages"
|
|
327
|
+
data = await self._post(url, payload, headers=headers)
|
|
328
|
+
|
|
329
|
+
# Extract content from Anthropic response format
|
|
330
|
+
content = data.get("content", [])
|
|
331
|
+
if content and isinstance(content, list):
|
|
332
|
+
return content[0].get("text", "")
|
|
333
|
+
return ""
|
|
334
|
+
|
|
335
|
+
# =========================================================================
|
|
336
|
+
# COMPLIANCE METHODS
|
|
337
|
+
# =========================================================================
|
|
338
|
+
|
|
339
|
+
async def check_compliance(
|
|
340
|
+
self,
|
|
341
|
+
content: str,
|
|
342
|
+
policy: str = "enterprise_standard",
|
|
343
|
+
) -> ComplianceResult:
|
|
344
|
+
"""
|
|
345
|
+
Check content for compliance issues.
|
|
346
|
+
|
|
347
|
+
Scans for:
|
|
348
|
+
- PII (emails, phone numbers, SSNs, etc.)
|
|
349
|
+
- Policy violations
|
|
350
|
+
- Regulatory issues (HIPAA, GDPR, SOC2)
|
|
351
|
+
|
|
352
|
+
Args:
|
|
353
|
+
content: Text to scan
|
|
354
|
+
policy: Compliance policy to apply
|
|
355
|
+
|
|
356
|
+
Returns:
|
|
357
|
+
ComplianceResult with violations, risk score, etc.
|
|
358
|
+
|
|
359
|
+
Example:
|
|
360
|
+
>>> result = await sf.check_compliance("My SSN is 123-45-6789")
|
|
361
|
+
>>> if result.has_violations:
|
|
362
|
+
... print(f"Found: {result.violations}")
|
|
363
|
+
"""
|
|
364
|
+
payload = {"content": content, "policy": policy}
|
|
365
|
+
url = f"{self.compliance_url}/api/compliance/scan"
|
|
366
|
+
data = await self._post(url, payload)
|
|
367
|
+
|
|
368
|
+
return ComplianceResult(
|
|
369
|
+
has_violations=not data.get("compliant", True),
|
|
370
|
+
compliance_score=100 - (data.get("risk_score", 0) * 100),
|
|
371
|
+
violations=data.get("violations", []),
|
|
372
|
+
pii_detected=data.get("pii_detected", []),
|
|
373
|
+
risk_level=self._risk_level_from_score(data.get("risk_score", 0)),
|
|
374
|
+
recommendations=data.get("recommendations", []),
|
|
375
|
+
redacted_content=data.get("redacted_content"),
|
|
376
|
+
)
|
|
377
|
+
|
|
378
|
+
def _risk_level_from_score(self, score: float) -> str:
|
|
379
|
+
"""Convert risk score to level."""
|
|
380
|
+
if score < 0.25:
|
|
381
|
+
return "low"
|
|
382
|
+
elif score < 0.5:
|
|
383
|
+
return "medium"
|
|
384
|
+
elif score < 0.75:
|
|
385
|
+
return "high"
|
|
386
|
+
return "critical"
|
|
387
|
+
|
|
388
|
+
async def redact_pii(self, content: str) -> str:
|
|
389
|
+
"""
|
|
390
|
+
Automatically redact PII from content.
|
|
391
|
+
|
|
392
|
+
Args:
|
|
393
|
+
content: Text potentially containing PII
|
|
394
|
+
|
|
395
|
+
Returns:
|
|
396
|
+
Redacted text
|
|
397
|
+
"""
|
|
398
|
+
result = await self.check_compliance(content)
|
|
399
|
+
return result.redacted_content or content
|
|
400
|
+
|
|
401
|
+
# =========================================================================
|
|
402
|
+
# INTELLIGENT COMPLIANCE (ML-POWERED)
|
|
403
|
+
# =========================================================================
|
|
404
|
+
|
|
405
|
+
async def intelligent_scan(
|
|
406
|
+
self,
|
|
407
|
+
content: str,
|
|
408
|
+
user_id: Optional[str] = None,
|
|
409
|
+
org_id: Optional[str] = None,
|
|
410
|
+
context: Optional[str] = None,
|
|
411
|
+
) -> IntelligentScanResult:
|
|
412
|
+
"""
|
|
413
|
+
Scan content using the ML-powered intelligent compliance engine.
|
|
414
|
+
|
|
415
|
+
This uses Smartflow's adaptive learning system which includes:
|
|
416
|
+
- Regex pattern matching (SSN, CC, email, phone, etc.)
|
|
417
|
+
- ML embedding similarity for semantic violation detection
|
|
418
|
+
- Behavioral analysis (user patterns, anomaly detection)
|
|
419
|
+
- Organization baselines (deviation from org norms)
|
|
420
|
+
|
|
421
|
+
Args:
|
|
422
|
+
content: Text to scan for compliance issues
|
|
423
|
+
user_id: Optional user ID for behavioral tracking
|
|
424
|
+
org_id: Optional organization ID for org-level baselines
|
|
425
|
+
context: Optional context (e.g., "customer_support", "sales")
|
|
426
|
+
|
|
427
|
+
Returns:
|
|
428
|
+
IntelligentScanResult with risk score, violations, and recommendations
|
|
429
|
+
|
|
430
|
+
Example:
|
|
431
|
+
>>> result = await sf.intelligent_scan(
|
|
432
|
+
... "My SSN is 123-45-6789",
|
|
433
|
+
... user_id="user123",
|
|
434
|
+
... org_id="acme_corp"
|
|
435
|
+
... )
|
|
436
|
+
>>> print(f"Risk: {result.risk_level}")
|
|
437
|
+
>>> print(f"Action: {result.recommended_action}")
|
|
438
|
+
"""
|
|
439
|
+
payload = {"content": content}
|
|
440
|
+
if user_id:
|
|
441
|
+
payload["user_id"] = user_id
|
|
442
|
+
if org_id:
|
|
443
|
+
payload["org_id"] = org_id
|
|
444
|
+
if context:
|
|
445
|
+
payload["context"] = context
|
|
446
|
+
|
|
447
|
+
url = f"{self.compliance_url}/api/compliance/intelligent/scan"
|
|
448
|
+
data = await self._post(url, payload)
|
|
449
|
+
return IntelligentScanResult.from_dict(data)
|
|
450
|
+
|
|
451
|
+
async def submit_compliance_feedback(
|
|
452
|
+
self,
|
|
453
|
+
scan_id: str,
|
|
454
|
+
is_false_positive: bool,
|
|
455
|
+
user_id: Optional[str] = None,
|
|
456
|
+
notes: Optional[str] = None,
|
|
457
|
+
) -> Dict[str, Any]:
|
|
458
|
+
"""
|
|
459
|
+
Submit feedback on a compliance scan result.
|
|
460
|
+
|
|
461
|
+
This feedback is used to train the ML model and reduce false positives.
|
|
462
|
+
|
|
463
|
+
Args:
|
|
464
|
+
scan_id: ID of the scan to provide feedback on
|
|
465
|
+
is_false_positive: True if the detection was a false positive
|
|
466
|
+
user_id: Optional user ID submitting feedback
|
|
467
|
+
notes: Optional notes explaining the feedback
|
|
468
|
+
|
|
469
|
+
Returns:
|
|
470
|
+
Confirmation response
|
|
471
|
+
|
|
472
|
+
Example:
|
|
473
|
+
>>> await sf.submit_compliance_feedback(
|
|
474
|
+
... scan_id="scan_abc123",
|
|
475
|
+
... is_false_positive=True,
|
|
476
|
+
... notes="This was a test phone number"
|
|
477
|
+
... )
|
|
478
|
+
"""
|
|
479
|
+
payload = {
|
|
480
|
+
"scan_id": scan_id,
|
|
481
|
+
"is_false_positive": is_false_positive,
|
|
482
|
+
}
|
|
483
|
+
if user_id:
|
|
484
|
+
payload["user_id"] = user_id
|
|
485
|
+
if notes:
|
|
486
|
+
payload["notes"] = notes
|
|
487
|
+
|
|
488
|
+
url = f"{self.compliance_url}/api/compliance/intelligent/feedback"
|
|
489
|
+
return await self._post(url, payload)
|
|
490
|
+
|
|
491
|
+
async def get_learning_status(self, user_id: str) -> LearningStatus:
|
|
492
|
+
"""
|
|
493
|
+
Get the learning status for a specific user.
|
|
494
|
+
|
|
495
|
+
Args:
|
|
496
|
+
user_id: User ID to check
|
|
497
|
+
|
|
498
|
+
Returns:
|
|
499
|
+
LearningStatus with progress info
|
|
500
|
+
"""
|
|
501
|
+
url = f"{self.compliance_url}/api/compliance/learning/status/{user_id}"
|
|
502
|
+
data = await self._get(url)
|
|
503
|
+
return LearningStatus.from_dict(data)
|
|
504
|
+
|
|
505
|
+
async def get_learning_summary(self) -> LearningSummary:
|
|
506
|
+
"""
|
|
507
|
+
Get overall learning summary across all users.
|
|
508
|
+
|
|
509
|
+
Returns:
|
|
510
|
+
LearningSummary with aggregate learning progress
|
|
511
|
+
"""
|
|
512
|
+
url = f"{self.compliance_url}/api/compliance/learning/summary"
|
|
513
|
+
data = await self._get(url)
|
|
514
|
+
return LearningSummary.from_dict(data)
|
|
515
|
+
|
|
516
|
+
async def get_ml_stats(self) -> MLStats:
|
|
517
|
+
"""
|
|
518
|
+
Get statistics about the ML compliance engine.
|
|
519
|
+
|
|
520
|
+
Returns:
|
|
521
|
+
MLStats with pattern counts and categories
|
|
522
|
+
"""
|
|
523
|
+
url = f"{self.compliance_url}/api/compliance/intelligent/stats"
|
|
524
|
+
data = await self._get(url)
|
|
525
|
+
ml_data = data.get("ml_stats", data)
|
|
526
|
+
return MLStats.from_dict(ml_data)
|
|
527
|
+
|
|
528
|
+
async def get_org_summary(self) -> Dict[str, Any]:
|
|
529
|
+
"""
|
|
530
|
+
Get organization-level compliance summary.
|
|
531
|
+
|
|
532
|
+
Returns:
|
|
533
|
+
Dict with org learning stats
|
|
534
|
+
"""
|
|
535
|
+
url = f"{self.compliance_url}/api/compliance/org/summary"
|
|
536
|
+
return await self._get(url)
|
|
537
|
+
|
|
538
|
+
async def get_org_baseline(self, org_id: str) -> OrgBaseline:
|
|
539
|
+
"""
|
|
540
|
+
Get the compliance baseline for a specific organization.
|
|
541
|
+
|
|
542
|
+
Args:
|
|
543
|
+
org_id: Organization ID
|
|
544
|
+
|
|
545
|
+
Returns:
|
|
546
|
+
OrgBaseline with org-level metrics
|
|
547
|
+
"""
|
|
548
|
+
url = f"{self.compliance_url}/api/compliance/org/status/{org_id}"
|
|
549
|
+
data = await self._get(url)
|
|
550
|
+
return OrgBaseline.from_dict(data)
|
|
551
|
+
|
|
552
|
+
async def get_persistence_stats(self) -> PersistenceStats:
|
|
553
|
+
"""
|
|
554
|
+
Get Redis persistence statistics for compliance data.
|
|
555
|
+
|
|
556
|
+
Returns:
|
|
557
|
+
PersistenceStats with storage info
|
|
558
|
+
"""
|
|
559
|
+
url = f"{self.compliance_url}/api/compliance/persistence/stats"
|
|
560
|
+
data = await self._get(url)
|
|
561
|
+
return PersistenceStats.from_dict(data.get("persistence", data))
|
|
562
|
+
|
|
563
|
+
async def save_compliance_data(self) -> Dict[str, Any]:
|
|
564
|
+
"""
|
|
565
|
+
Trigger manual save of compliance data to Redis.
|
|
566
|
+
|
|
567
|
+
Returns:
|
|
568
|
+
Confirmation response
|
|
569
|
+
"""
|
|
570
|
+
url = f"{self.compliance_url}/api/compliance/persistence/save"
|
|
571
|
+
return await self._post(url, {})
|
|
572
|
+
|
|
573
|
+
async def get_intelligent_health(self) -> Dict[str, Any]:
|
|
574
|
+
"""
|
|
575
|
+
Get health status of the intelligent compliance engine.
|
|
576
|
+
|
|
577
|
+
Returns:
|
|
578
|
+
Health status including ML engine, behavior analysis, etc.
|
|
579
|
+
"""
|
|
580
|
+
url = f"{self.compliance_url}/api/compliance/intelligent/health"
|
|
581
|
+
return await self._get(url)
|
|
582
|
+
|
|
583
|
+
# =========================================================================
|
|
584
|
+
# CACHE METHODS
|
|
585
|
+
# =========================================================================
|
|
586
|
+
|
|
587
|
+
async def get_cache_stats(self) -> CacheStats:
|
|
588
|
+
"""
|
|
589
|
+
Get Smartflow cache statistics.
|
|
590
|
+
|
|
591
|
+
Returns L1/L2/L3 hit rates, tokens saved, cost savings, etc.
|
|
592
|
+
|
|
593
|
+
Returns:
|
|
594
|
+
CacheStats object with cache metrics
|
|
595
|
+
|
|
596
|
+
Example:
|
|
597
|
+
>>> stats = await sf.get_cache_stats()
|
|
598
|
+
>>> print(f"Hit rate: {stats.hit_rate:.1%}")
|
|
599
|
+
>>> print(f"Tokens saved: {stats.tokens_saved:,}")
|
|
600
|
+
"""
|
|
601
|
+
url = f"{self.management_url}/api/metacache/stats"
|
|
602
|
+
data = await self._get(url)
|
|
603
|
+
|
|
604
|
+
# Handle wrapped response
|
|
605
|
+
if "data" in data:
|
|
606
|
+
data = data["data"]
|
|
607
|
+
|
|
608
|
+
return CacheStats.from_dict(data)
|
|
609
|
+
|
|
610
|
+
# =========================================================================
|
|
611
|
+
# HEALTH & MONITORING
|
|
612
|
+
# =========================================================================
|
|
613
|
+
|
|
614
|
+
async def health(self) -> Dict[str, Any]:
|
|
615
|
+
"""
|
|
616
|
+
Quick health check of Smartflow proxy.
|
|
617
|
+
|
|
618
|
+
Returns:
|
|
619
|
+
Health status dict
|
|
620
|
+
"""
|
|
621
|
+
url = f"{self.base_url}/health"
|
|
622
|
+
return await self._get(url)
|
|
623
|
+
|
|
624
|
+
async def health_comprehensive(self) -> SystemHealth:
|
|
625
|
+
"""
|
|
626
|
+
Comprehensive health check including all services and providers.
|
|
627
|
+
|
|
628
|
+
Returns:
|
|
629
|
+
SystemHealth object with full status
|
|
630
|
+
"""
|
|
631
|
+
url = f"{self.management_url}/api/health/comprehensive"
|
|
632
|
+
data = await self._get(url)
|
|
633
|
+
|
|
634
|
+
if "data" in data:
|
|
635
|
+
data = data["data"]
|
|
636
|
+
|
|
637
|
+
return SystemHealth.from_dict(data)
|
|
638
|
+
|
|
639
|
+
async def get_provider_health(self) -> List[ProviderHealth]:
|
|
640
|
+
"""
|
|
641
|
+
Get health status of all AI providers.
|
|
642
|
+
|
|
643
|
+
Returns:
|
|
644
|
+
List of ProviderHealth objects
|
|
645
|
+
"""
|
|
646
|
+
url = f"{self.management_url}/api/providers/perf"
|
|
647
|
+
data = await self._get(url)
|
|
648
|
+
|
|
649
|
+
snapshots = data.get("data", {}).get("snapshots", [])
|
|
650
|
+
return [ProviderHealth.from_dict(s) for s in snapshots]
|
|
651
|
+
|
|
652
|
+
# =========================================================================
|
|
653
|
+
# VAS LOGS (AUDIT)
|
|
654
|
+
# =========================================================================
|
|
655
|
+
|
|
656
|
+
async def get_logs(
|
|
657
|
+
self,
|
|
658
|
+
limit: int = 50,
|
|
659
|
+
provider: Optional[str] = None,
|
|
660
|
+
) -> List[VASLog]:
|
|
661
|
+
"""
|
|
662
|
+
Get VAS audit logs.
|
|
663
|
+
|
|
664
|
+
Full audit trail of all AI interactions through Smartflow.
|
|
665
|
+
|
|
666
|
+
Args:
|
|
667
|
+
limit: Maximum logs to return
|
|
668
|
+
provider: Filter by provider name
|
|
669
|
+
|
|
670
|
+
Returns:
|
|
671
|
+
List of VASLog objects
|
|
672
|
+
"""
|
|
673
|
+
url = f"{self.management_url}/api/vas/logs"
|
|
674
|
+
params = {"limit": limit}
|
|
675
|
+
if provider:
|
|
676
|
+
params["provider"] = provider
|
|
677
|
+
|
|
678
|
+
data = await self._get(url, params=params)
|
|
679
|
+
|
|
680
|
+
logs_data = data.get("data", [])
|
|
681
|
+
return [VASLog.from_dict(log) for log in logs_data]
|
|
682
|
+
|
|
683
|
+
async def get_logs_hybrid(
|
|
684
|
+
self,
|
|
685
|
+
limit: int = 100,
|
|
686
|
+
) -> List[Dict[str, Any]]:
|
|
687
|
+
"""
|
|
688
|
+
Get VAS logs from hybrid bridge (Redis + MongoDB combined).
|
|
689
|
+
|
|
690
|
+
Args:
|
|
691
|
+
limit: Maximum logs to return
|
|
692
|
+
|
|
693
|
+
Returns:
|
|
694
|
+
List of log dicts
|
|
695
|
+
"""
|
|
696
|
+
url = f"{self.bridge_url}/api/redis/logs"
|
|
697
|
+
params = {"limit": limit}
|
|
698
|
+
data = await self._get(url, params=params)
|
|
699
|
+
return data.get("logs", [])
|
|
700
|
+
|
|
701
|
+
# =========================================================================
|
|
702
|
+
# ANALYTICS
|
|
703
|
+
# =========================================================================
|
|
704
|
+
|
|
705
|
+
async def get_analytics(
|
|
706
|
+
self,
|
|
707
|
+
start_date: Optional[str] = None,
|
|
708
|
+
end_date: Optional[str] = None,
|
|
709
|
+
) -> Dict[str, Any]:
|
|
710
|
+
"""
|
|
711
|
+
Get usage analytics.
|
|
712
|
+
|
|
713
|
+
Args:
|
|
714
|
+
start_date: ISO format start date
|
|
715
|
+
end_date: ISO format end date
|
|
716
|
+
|
|
717
|
+
Returns:
|
|
718
|
+
Analytics data dict
|
|
719
|
+
"""
|
|
720
|
+
url = f"{self.bridge_url}/api/hybrid/analytics"
|
|
721
|
+
params = {}
|
|
722
|
+
if start_date:
|
|
723
|
+
params["start_date"] = start_date
|
|
724
|
+
if end_date:
|
|
725
|
+
params["end_date"] = end_date
|
|
726
|
+
|
|
727
|
+
return await self._get(url, params=params)
|
|
728
|
+
|
|
729
|
+
# =========================================================================
|
|
730
|
+
# ROUTING
|
|
731
|
+
# =========================================================================
|
|
732
|
+
|
|
733
|
+
async def get_routing_status(self) -> Dict[str, Any]:
|
|
734
|
+
"""
|
|
735
|
+
Get current routing configuration and status.
|
|
736
|
+
|
|
737
|
+
Returns:
|
|
738
|
+
Routing status including active providers, failover state, etc.
|
|
739
|
+
"""
|
|
740
|
+
url = f"{self.management_url}/api/routing/status"
|
|
741
|
+
data = await self._get(url)
|
|
742
|
+
return data.get("data", data)
|
|
743
|
+
|
|
744
|
+
async def force_provider(
|
|
745
|
+
self,
|
|
746
|
+
provider: str,
|
|
747
|
+
duration_seconds: int = 300,
|
|
748
|
+
) -> Dict[str, Any]:
|
|
749
|
+
"""
|
|
750
|
+
Force routing to a specific provider.
|
|
751
|
+
|
|
752
|
+
Args:
|
|
753
|
+
provider: Provider name ("openai", "anthropic", etc.)
|
|
754
|
+
duration_seconds: How long to force (default 5 minutes)
|
|
755
|
+
|
|
756
|
+
Returns:
|
|
757
|
+
Confirmation response
|
|
758
|
+
"""
|
|
759
|
+
url = f"{self.management_url}/api/routing/override"
|
|
760
|
+
payload = {
|
|
761
|
+
"provider": provider,
|
|
762
|
+
"ttl_minutes": duration_seconds // 60,
|
|
763
|
+
}
|
|
764
|
+
return await self._post(url, payload)
|
|
765
|
+
|
|
766
|
+
# =========================================================================
|
|
767
|
+
# CHATBOT (Built-in Smartflow chatbot)
|
|
768
|
+
# =========================================================================
|
|
769
|
+
|
|
770
|
+
async def chatbot_query(self, query: str) -> Dict[str, Any]:
|
|
771
|
+
"""
|
|
772
|
+
Query Smartflow's built-in chatbot for system info.
|
|
773
|
+
|
|
774
|
+
The chatbot can answer questions about:
|
|
775
|
+
- VAS logs and analytics
|
|
776
|
+
- Cache performance
|
|
777
|
+
- System health
|
|
778
|
+
- Cost analysis
|
|
779
|
+
|
|
780
|
+
Args:
|
|
781
|
+
query: Natural language query (e.g., "show cache stats")
|
|
782
|
+
|
|
783
|
+
Returns:
|
|
784
|
+
Chatbot response
|
|
785
|
+
|
|
786
|
+
Example:
|
|
787
|
+
>>> result = await sf.chatbot_query("show me today's cache stats")
|
|
788
|
+
>>> print(result["response"])
|
|
789
|
+
"""
|
|
790
|
+
url = f"{self.base_url}/api/chatbot/query"
|
|
791
|
+
payload = {"query": query}
|
|
792
|
+
return await self._post(url, payload)
|
|
793
|
+
|
|
794
|
+
# =========================================================================
|
|
795
|
+
# CONTEXT MANAGER
|
|
796
|
+
# =========================================================================
|
|
797
|
+
|
|
798
|
+
async def close(self):
|
|
799
|
+
"""Close the HTTP client."""
|
|
800
|
+
if self._client:
|
|
801
|
+
await self._client.aclose()
|
|
802
|
+
self._client = None
|
|
803
|
+
|
|
804
|
+
async def __aenter__(self):
|
|
805
|
+
"""Async context manager entry."""
|
|
806
|
+
await self._ensure_client()
|
|
807
|
+
return self
|
|
808
|
+
|
|
809
|
+
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
|
810
|
+
"""Async context manager exit."""
|
|
811
|
+
await self.close()
|
|
812
|
+
|
|
813
|
+
|
|
814
|
+
# =============================================================================
|
|
815
|
+
# AGENT BUILDER
|
|
816
|
+
# =============================================================================
|
|
817
|
+
|
|
818
|
+
class SmartflowAgent:
|
|
819
|
+
"""
|
|
820
|
+
A Smartflow-powered AI agent with built-in compliance, caching, and routing.
|
|
821
|
+
|
|
822
|
+
SmartflowAgent wraps a SmartflowClient to provide higher-level abstractions
|
|
823
|
+
for building AI applications, including:
|
|
824
|
+
- Conversation memory
|
|
825
|
+
- Automatic compliance scanning
|
|
826
|
+
- Tool/function calling
|
|
827
|
+
- Response caching
|
|
828
|
+
|
|
829
|
+
Example:
|
|
830
|
+
>>> async with SmartflowClient("http://smartflow:7775") as sf:
|
|
831
|
+
... agent = SmartflowAgent(
|
|
832
|
+
... client=sf,
|
|
833
|
+
... name="CustomerSupport",
|
|
834
|
+
... system_prompt="You are a helpful customer support agent.",
|
|
835
|
+
... compliance_policy="enterprise_standard"
|
|
836
|
+
... )
|
|
837
|
+
...
|
|
838
|
+
... response = await agent.chat("How do I reset my password?")
|
|
839
|
+
... print(response)
|
|
840
|
+
"""
|
|
841
|
+
|
|
842
|
+
def __init__(
|
|
843
|
+
self,
|
|
844
|
+
client: SmartflowClient,
|
|
845
|
+
name: str = "SmartflowAgent",
|
|
846
|
+
model: str = "gpt-4o",
|
|
847
|
+
system_prompt: Optional[str] = None,
|
|
848
|
+
temperature: float = 0.7,
|
|
849
|
+
max_tokens: Optional[int] = None,
|
|
850
|
+
compliance_policy: str = "enterprise_standard",
|
|
851
|
+
enable_compliance_scan: bool = True,
|
|
852
|
+
user_id: Optional[str] = None,
|
|
853
|
+
org_id: Optional[str] = None,
|
|
854
|
+
tools: Optional[List[Dict[str, Any]]] = None,
|
|
855
|
+
):
|
|
856
|
+
"""
|
|
857
|
+
Initialize a SmartflowAgent.
|
|
858
|
+
|
|
859
|
+
Args:
|
|
860
|
+
client: SmartflowClient instance
|
|
861
|
+
name: Agent name for logging
|
|
862
|
+
model: AI model to use
|
|
863
|
+
system_prompt: System prompt for the agent
|
|
864
|
+
temperature: Sampling temperature
|
|
865
|
+
max_tokens: Maximum tokens per response
|
|
866
|
+
compliance_policy: Compliance policy to apply
|
|
867
|
+
enable_compliance_scan: Scan inputs/outputs for compliance
|
|
868
|
+
user_id: User ID for behavioral tracking
|
|
869
|
+
org_id: Organization ID for org baselines
|
|
870
|
+
tools: List of tools/functions the agent can call
|
|
871
|
+
"""
|
|
872
|
+
self.client = client
|
|
873
|
+
self.name = name
|
|
874
|
+
self.model = model
|
|
875
|
+
self.system_prompt = system_prompt
|
|
876
|
+
self.temperature = temperature
|
|
877
|
+
self.max_tokens = max_tokens
|
|
878
|
+
self.compliance_policy = compliance_policy
|
|
879
|
+
self.enable_compliance_scan = enable_compliance_scan
|
|
880
|
+
self.user_id = user_id
|
|
881
|
+
self.org_id = org_id
|
|
882
|
+
self.tools = tools or []
|
|
883
|
+
|
|
884
|
+
# Conversation memory
|
|
885
|
+
self._messages: List[Dict[str, str]] = []
|
|
886
|
+
if system_prompt:
|
|
887
|
+
self._messages.append({"role": "system", "content": system_prompt})
|
|
888
|
+
|
|
889
|
+
async def chat(
|
|
890
|
+
self,
|
|
891
|
+
message: str,
|
|
892
|
+
scan_input: bool = True,
|
|
893
|
+
scan_output: bool = True,
|
|
894
|
+
) -> str:
|
|
895
|
+
"""
|
|
896
|
+
Send a message to the agent.
|
|
897
|
+
|
|
898
|
+
Args:
|
|
899
|
+
message: User message
|
|
900
|
+
scan_input: Scan input for compliance (default True)
|
|
901
|
+
scan_output: Scan output for compliance (default True)
|
|
902
|
+
|
|
903
|
+
Returns:
|
|
904
|
+
Agent's response
|
|
905
|
+
|
|
906
|
+
Raises:
|
|
907
|
+
SmartflowError: If compliance violation blocks the message
|
|
908
|
+
"""
|
|
909
|
+
# Optionally scan input
|
|
910
|
+
if self.enable_compliance_scan and scan_input:
|
|
911
|
+
input_scan = await self.client.intelligent_scan(
|
|
912
|
+
content=message,
|
|
913
|
+
user_id=self.user_id,
|
|
914
|
+
org_id=self.org_id,
|
|
915
|
+
)
|
|
916
|
+
if input_scan.recommended_action == "Block":
|
|
917
|
+
raise SmartflowError(
|
|
918
|
+
f"Input blocked by compliance: {input_scan.explanation}"
|
|
919
|
+
)
|
|
920
|
+
|
|
921
|
+
# Add message to history
|
|
922
|
+
self._messages.append({"role": "user", "content": message})
|
|
923
|
+
|
|
924
|
+
# Get response
|
|
925
|
+
response = await self.client.chat_completions(
|
|
926
|
+
messages=self._messages,
|
|
927
|
+
model=self.model,
|
|
928
|
+
temperature=self.temperature,
|
|
929
|
+
max_tokens=self.max_tokens,
|
|
930
|
+
)
|
|
931
|
+
|
|
932
|
+
assistant_message = response.content
|
|
933
|
+
|
|
934
|
+
# Optionally scan output
|
|
935
|
+
if self.enable_compliance_scan and scan_output:
|
|
936
|
+
output_scan = await self.client.intelligent_scan(
|
|
937
|
+
content=assistant_message,
|
|
938
|
+
user_id=self.user_id,
|
|
939
|
+
org_id=self.org_id,
|
|
940
|
+
)
|
|
941
|
+
if output_scan.recommended_action == "Block":
|
|
942
|
+
# Don't add to history, return warning
|
|
943
|
+
return "[Response blocked due to compliance policy]"
|
|
944
|
+
|
|
945
|
+
# Add response to history
|
|
946
|
+
self._messages.append({"role": "assistant", "content": assistant_message})
|
|
947
|
+
|
|
948
|
+
return assistant_message
|
|
949
|
+
|
|
950
|
+
def clear_history(self):
|
|
951
|
+
"""Clear conversation history, keeping system prompt."""
|
|
952
|
+
self._messages = []
|
|
953
|
+
if self.system_prompt:
|
|
954
|
+
self._messages.append({"role": "system", "content": self.system_prompt})
|
|
955
|
+
|
|
956
|
+
def get_history(self) -> List[Dict[str, str]]:
|
|
957
|
+
"""Get conversation history."""
|
|
958
|
+
return self._messages.copy()
|
|
959
|
+
|
|
960
|
+
@property
|
|
961
|
+
def message_count(self) -> int:
|
|
962
|
+
"""Get number of messages in history."""
|
|
963
|
+
return len(self._messages)
|
|
964
|
+
|
|
965
|
+
|
|
966
|
+
class SmartflowWorkflow:
|
|
967
|
+
"""
|
|
968
|
+
A workflow orchestrator for chaining AI operations.
|
|
969
|
+
|
|
970
|
+
SmartflowWorkflow allows you to define and execute multi-step AI workflows
|
|
971
|
+
with branching, parallel execution, and error handling.
|
|
972
|
+
|
|
973
|
+
Example:
|
|
974
|
+
>>> workflow = SmartflowWorkflow(client, name="SupportTicketFlow")
|
|
975
|
+
>>>
|
|
976
|
+
>>> workflow.add_step("classify", action="chat",
|
|
977
|
+
... config={"prompt": "Classify this ticket: {input}"})
|
|
978
|
+
>>> workflow.add_step("route", action="condition",
|
|
979
|
+
... config={"field": "category", "cases": {...}})
|
|
980
|
+
>>>
|
|
981
|
+
>>> result = await workflow.execute({"input": ticket_text})
|
|
982
|
+
"""
|
|
983
|
+
|
|
984
|
+
def __init__(
|
|
985
|
+
self,
|
|
986
|
+
client: SmartflowClient,
|
|
987
|
+
name: str = "SmartflowWorkflow",
|
|
988
|
+
):
|
|
989
|
+
"""
|
|
990
|
+
Initialize a workflow.
|
|
991
|
+
|
|
992
|
+
Args:
|
|
993
|
+
client: SmartflowClient instance
|
|
994
|
+
name: Workflow name for logging
|
|
995
|
+
"""
|
|
996
|
+
self.client = client
|
|
997
|
+
self.name = name
|
|
998
|
+
self.steps: Dict[str, WorkflowStep] = {}
|
|
999
|
+
self.entry_step: Optional[str] = None
|
|
1000
|
+
|
|
1001
|
+
def add_step(
|
|
1002
|
+
self,
|
|
1003
|
+
name: str,
|
|
1004
|
+
action: str,
|
|
1005
|
+
config: Dict[str, Any] = None,
|
|
1006
|
+
next_steps: List[str] = None,
|
|
1007
|
+
on_error: Optional[str] = None,
|
|
1008
|
+
) -> "SmartflowWorkflow":
|
|
1009
|
+
"""
|
|
1010
|
+
Add a step to the workflow.
|
|
1011
|
+
|
|
1012
|
+
Args:
|
|
1013
|
+
name: Step name
|
|
1014
|
+
action: Action type ("chat", "compliance_check", "condition", etc.)
|
|
1015
|
+
config: Step configuration
|
|
1016
|
+
next_steps: Names of subsequent steps
|
|
1017
|
+
on_error: Step to execute on error
|
|
1018
|
+
|
|
1019
|
+
Returns:
|
|
1020
|
+
Self for chaining
|
|
1021
|
+
"""
|
|
1022
|
+
step = WorkflowStep(
|
|
1023
|
+
name=name,
|
|
1024
|
+
action=action,
|
|
1025
|
+
config=config or {},
|
|
1026
|
+
next_steps=next_steps or [],
|
|
1027
|
+
on_error=on_error,
|
|
1028
|
+
)
|
|
1029
|
+
self.steps[name] = step
|
|
1030
|
+
|
|
1031
|
+
# First step added becomes entry
|
|
1032
|
+
if self.entry_step is None:
|
|
1033
|
+
self.entry_step = name
|
|
1034
|
+
|
|
1035
|
+
return self
|
|
1036
|
+
|
|
1037
|
+
def set_entry(self, step_name: str) -> "SmartflowWorkflow":
|
|
1038
|
+
"""Set the entry step for the workflow."""
|
|
1039
|
+
if step_name not in self.steps:
|
|
1040
|
+
raise ValueError(f"Step '{step_name}' not found in workflow")
|
|
1041
|
+
self.entry_step = step_name
|
|
1042
|
+
return self
|
|
1043
|
+
|
|
1044
|
+
async def execute(
|
|
1045
|
+
self,
|
|
1046
|
+
input_data: Dict[str, Any],
|
|
1047
|
+
max_iterations: int = 100,
|
|
1048
|
+
) -> WorkflowResult:
|
|
1049
|
+
"""
|
|
1050
|
+
Execute the workflow.
|
|
1051
|
+
|
|
1052
|
+
Args:
|
|
1053
|
+
input_data: Initial input data
|
|
1054
|
+
max_iterations: Maximum steps to execute (prevent infinite loops)
|
|
1055
|
+
|
|
1056
|
+
Returns:
|
|
1057
|
+
WorkflowResult with output and execution details
|
|
1058
|
+
"""
|
|
1059
|
+
if not self.entry_step:
|
|
1060
|
+
return WorkflowResult(
|
|
1061
|
+
success=False,
|
|
1062
|
+
output=None,
|
|
1063
|
+
errors=["No entry step defined"],
|
|
1064
|
+
)
|
|
1065
|
+
|
|
1066
|
+
context = {"input": input_data, "output": None}
|
|
1067
|
+
steps_executed = []
|
|
1068
|
+
errors = []
|
|
1069
|
+
current_step = self.entry_step
|
|
1070
|
+
iterations = 0
|
|
1071
|
+
total_tokens = 0
|
|
1072
|
+
|
|
1073
|
+
import time
|
|
1074
|
+
start_time = time.time()
|
|
1075
|
+
|
|
1076
|
+
while current_step and iterations < max_iterations:
|
|
1077
|
+
iterations += 1
|
|
1078
|
+
step = self.steps.get(current_step)
|
|
1079
|
+
|
|
1080
|
+
if not step:
|
|
1081
|
+
errors.append(f"Step '{current_step}' not found")
|
|
1082
|
+
break
|
|
1083
|
+
|
|
1084
|
+
steps_executed.append(current_step)
|
|
1085
|
+
|
|
1086
|
+
try:
|
|
1087
|
+
result, next_step = await self._execute_step(step, context)
|
|
1088
|
+
context["output"] = result
|
|
1089
|
+
current_step = next_step
|
|
1090
|
+
except Exception as e:
|
|
1091
|
+
errors.append(f"Error in step '{current_step}': {str(e)}")
|
|
1092
|
+
if step.on_error:
|
|
1093
|
+
current_step = step.on_error
|
|
1094
|
+
else:
|
|
1095
|
+
break
|
|
1096
|
+
|
|
1097
|
+
execution_time_ms = (time.time() - start_time) * 1000
|
|
1098
|
+
|
|
1099
|
+
return WorkflowResult(
|
|
1100
|
+
success=len(errors) == 0,
|
|
1101
|
+
output=context.get("output"),
|
|
1102
|
+
steps_executed=steps_executed,
|
|
1103
|
+
errors=errors,
|
|
1104
|
+
total_tokens=total_tokens,
|
|
1105
|
+
execution_time_ms=execution_time_ms,
|
|
1106
|
+
)
|
|
1107
|
+
|
|
1108
|
+
async def _execute_step(
|
|
1109
|
+
self,
|
|
1110
|
+
step: WorkflowStep,
|
|
1111
|
+
context: Dict[str, Any],
|
|
1112
|
+
) -> tuple:
|
|
1113
|
+
"""Execute a single workflow step."""
|
|
1114
|
+
action = step.action
|
|
1115
|
+
config = step.config
|
|
1116
|
+
|
|
1117
|
+
if action == "chat":
|
|
1118
|
+
prompt = config.get("prompt", "{input}")
|
|
1119
|
+
formatted_prompt = self._format_template(prompt, context)
|
|
1120
|
+
|
|
1121
|
+
response = await self.client.chat(
|
|
1122
|
+
message=formatted_prompt,
|
|
1123
|
+
model=config.get("model", "gpt-4o"),
|
|
1124
|
+
temperature=config.get("temperature", 0.7),
|
|
1125
|
+
)
|
|
1126
|
+
|
|
1127
|
+
next_step = step.next_steps[0] if step.next_steps else None
|
|
1128
|
+
return response, next_step
|
|
1129
|
+
|
|
1130
|
+
elif action == "compliance_check":
|
|
1131
|
+
content = config.get("content", context.get("output", ""))
|
|
1132
|
+
if isinstance(content, str):
|
|
1133
|
+
content = self._format_template(content, context)
|
|
1134
|
+
|
|
1135
|
+
result = await self.client.intelligent_scan(content=content)
|
|
1136
|
+
|
|
1137
|
+
next_step = step.next_steps[0] if step.next_steps else None
|
|
1138
|
+
return result, next_step
|
|
1139
|
+
|
|
1140
|
+
elif action == "condition":
|
|
1141
|
+
field = config.get("field", "output")
|
|
1142
|
+
value = context.get(field)
|
|
1143
|
+
cases = config.get("cases", {})
|
|
1144
|
+
default = config.get("default")
|
|
1145
|
+
|
|
1146
|
+
next_step = cases.get(str(value), default)
|
|
1147
|
+
return value, next_step
|
|
1148
|
+
|
|
1149
|
+
else:
|
|
1150
|
+
raise ValueError(f"Unknown action: {action}")
|
|
1151
|
+
|
|
1152
|
+
def _format_template(self, template: str, context: Dict[str, Any]) -> str:
|
|
1153
|
+
"""Format a template string with context values."""
|
|
1154
|
+
result = template
|
|
1155
|
+
for key, value in context.items():
|
|
1156
|
+
placeholder = "{" + key + "}"
|
|
1157
|
+
if placeholder in result:
|
|
1158
|
+
result = result.replace(placeholder, str(value))
|
|
1159
|
+
return result
|
|
1160
|
+
|