File size: 8,950 Bytes
0e3d4b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
"""REST API Client — connect to external APIs and services.

Features:
- Generic REST client (GET, POST, PUT, DELETE)
- Authentication: API key, Bearer token, Basic auth, custom headers
- Request/response logging
- Rate limiting (configurable requests per second)
- Retry with exponential backoff
- Timeout handling
- JSON and raw response support
- Connection pooling (via urllib)

Uses only stdlib urllib — no external dependencies.
100% local: no data sent to any cloud service unless explicitly configured.
"""

from __future__ import annotations

import json
import logging
import time
import urllib.error
import urllib.request
from dataclasses import dataclass, field
from typing import Any

logger = logging.getLogger(__name__)


@dataclass
class APIConfig:
    """Configuration for an API connection."""
    name: str
    base_url: str
    api_key: str = ""
    auth_type: str = "api_key"  # "api_key", "bearer", "basic", "none", "custom"
    headers: dict[str, str] = field(default_factory=dict)
    timeout_s: float = 30.0
    max_retries: int = 3
    rate_limit_s: float = 0.0  # min seconds between requests
    retry_backoff: float = 1.5


@dataclass
class APIResponse:
    """Response from an API call."""
    success: bool
    status_code: int
    data: Any = None
    error: str = ""
    elapsed_s: float = 0.0
    url: str = ""


class RESTClient:
    """Generic REST API client with auth, retries, and rate limiting."""

    def __init__(self, config: APIConfig) -> None:
        self.config = config
        self._last_request_time = 0.0
        self._stats = {
            "total_requests": 0,
            "successful_requests": 0,
            "failed_requests": 0,
            "retries": 0,
            "avg_response_time_s": 0.0,
        }

    def _build_headers(self, extra: dict[str, str] | None = None) -> dict[str, str]:
        """Build request headers with authentication."""
        headers = dict(self.config.headers)
        if extra:
            headers.update(extra)

        if self.config.auth_type == "api_key" and self.config.api_key:
            headers["X-API-Key"] = self.config.api_key
        elif self.config.auth_type == "bearer" and self.config.api_key:
            headers["Authorization"] = f"Bearer {self.config.api_key}"
        elif self.config.auth_type == "basic" and self.config.api_key:
            import base64
            headers["Authorization"] = f"Basic {base64.b64encode(self.config.api_key.encode()).decode()}"

        return headers

    def _rate_limit(self) -> None:
        """Enforce rate limiting."""
        if self.config.rate_limit_s > 0:
            elapsed = time.time() - self._last_request_time
            if elapsed < self.config.rate_limit_s:
                time.sleep(self.config.rate_limit_s - elapsed)
        self._last_request_time = time.time()

    def request(self, method: str, endpoint: str, data: dict | None = None,
                params: dict | None = None, headers: dict | None = None) -> APIResponse:
        """Make an HTTP request.

        Args:
            method: GET, POST, PUT, DELETE
            endpoint: API endpoint (appended to base_url)
            data: request body (JSON)
            params: query parameters
            headers: extra headers
        """
        url = self._build_url(endpoint, params)
        body = json.dumps(data).encode() if data else None
        req_headers = self._build_headers(headers)
        if body:
            req_headers["Content-Type"] = "application/json"

        for attempt in range(self.config.max_retries + 1):
            self._rate_limit()
            t0 = time.time()
            self._stats["total_requests"] += 1

            try:
                req = urllib.request.Request(url, data=body, method=method, headers=req_headers)
                with urllib.request.urlopen(req, timeout=self.config.timeout_s) as resp:
                    raw = resp.read().decode()
                    elapsed = time.time() - t0
                    self._stats["successful_requests"] += 1
                    self._update_avg_time(elapsed)

                    try:
                        parsed = json.loads(raw)
                    except json.JSONDecodeError:
                        parsed = raw

                    return APIResponse(
                        success=True, status_code=resp.status,
                        data=parsed, elapsed_s=elapsed, url=url,
                    )

            except urllib.error.HTTPError as e:
                elapsed = time.time() - t0
                error_body = ""
                try:
                    error_body = e.read().decode()
                except Exception:
                    pass

                if attempt < self.config.max_retries and e.code >= 500:
                    self._stats["retries"] += 1
                    wait = self.config.retry_backoff ** (attempt + 1)
                    logger.warning("Retry %d/%d for %s (HTTP %d) after %.1fs",
                                   attempt + 1, self.config.max_retries, url, e.code, wait)
                    time.sleep(wait)
                    continue

                self._stats["failed_requests"] += 1
                return APIResponse(
                    success=False, status_code=e.code,
                    error=f"HTTP {e.code}: {error_body[:200]}", elapsed_s=elapsed, url=url,
                )

            except Exception as e:
                elapsed = time.time() - t0
                if attempt < self.config.max_retries:
                    self._stats["retries"] += 1
                    wait = self.config.retry_backoff ** (attempt + 1)
                    logger.warning("Retry %d/%d for %s: %s", attempt + 1, self.config.max_retries, url, e)
                    time.sleep(wait)
                    continue

                self._stats["failed_requests"] += 1
                return APIResponse(
                    success=False, status_code=0,
                    error=str(e), elapsed_s=elapsed, url=url,
                )

        return APIResponse(success=False, status_code=0, error="Max retries exceeded", url=url)

    def get(self, endpoint: str, params: dict | None = None) -> APIResponse:
        return self.request("GET", endpoint, params=params)

    def post(self, endpoint: str, data: dict | None = None) -> APIResponse:
        return self.request("POST", endpoint, data=data)

    def put(self, endpoint: str, data: dict | None = None) -> APIResponse:
        return self.request("PUT", endpoint, data=data)

    def delete(self, endpoint: str) -> APIResponse:
        return self.request("DELETE", endpoint)

    def _build_url(self, endpoint: str, params: dict | None = None) -> str:
        url = f"{self.config.base_url.rstrip('/')}/{endpoint.lstrip('/')}"
        if params:
            import urllib.parse
            query = urllib.parse.urlencode(params)
            url = f"{url}?{query}"
        return url

    def _update_avg_time(self, elapsed: float) -> None:
        total = self._stats["successful_requests"]
        self._stats["avg_response_time_s"] = (
            (self._stats["avg_response_time_s"] * (total - 1) + elapsed) / total
        )

    def get_stats(self) -> dict[str, Any]:
        return {**self._stats, "name": self.config.name, "base_url": self.config.base_url}


class ConnectorRegistry:
    """Registry of named API connectors.

    Allows the LLM to connect to multiple external services.
    Connectors are registered with a name and can be called by the LLM
    via tool calls: [TOOL: api_call("service_name", "GET", "/endpoint")]
    """

    def __init__(self) -> None:
        self._connectors: dict[str, RESTClient] = {}
        self._stats = {"total_connectors": 0, "total_calls": 0}

    def register(self, name: str, config: APIConfig) -> None:
        """Register a named API connector."""
        config.name = name
        client = RESTClient(config)
        self._connectors[name] = client
        self._stats["total_connectors"] += 1
        logger.info("Registered API connector: %s → %s", name, config.base_url)

    def get(self, name: str) -> RESTClient | None:
        return self._connectors.get(name)

    def call(self, name: str, method: str, endpoint: str,
             data: dict | None = None, params: dict | None = None) -> APIResponse:
        """Call a registered connector."""
        client = self._connectors.get(name)
        if client is None:
            return APIResponse(success=False, status_code=0, error=f"Connector '{name}' not found")
        self._stats["total_calls"] += 1
        return client.request(method, endpoint, data=data, params=params)

    def list_connectors(self) -> list[dict[str, Any]]:
        return [{"name": name, **client.get_stats()} for name, client in self._connectors.items()]

    def get_stats(self) -> dict[str, Any]:
        return {**self._stats, "connectors": {name: c.get_stats() for name, c in self._connectors.items()}}