singlestore-auth-iam 0.4.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.
- s2iam/__init__.py +35 -0
- s2iam/api.py +300 -0
- s2iam/aws/__init__.py +271 -0
- s2iam/azure/__init__.py +456 -0
- s2iam/gcp/__init__.py +428 -0
- s2iam/https.py +16 -0
- s2iam/jwt.py +203 -0
- s2iam/models.py +129 -0
- singlestore_auth_iam-0.4.0.dist-info/METADATA +270 -0
- singlestore_auth_iam-0.4.0.dist-info/RECORD +11 -0
- singlestore_auth_iam-0.4.0.dist-info/WHEEL +4 -0
s2iam/__init__.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""
|
|
2
|
+
SingleStore Auth IAM - Python Client Library
|
|
3
|
+
|
|
4
|
+
A Python client library for cloud provider identity detection and authentication
|
|
5
|
+
with SingleStore's IAM service.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
__version__ = "0.4.0"
|
|
9
|
+
|
|
10
|
+
from .api import DETECT_PROVIDER_DEFAULT_TIMEOUT, detect_provider
|
|
11
|
+
from .jwt import get_jwt, get_jwt_api, get_jwt_database
|
|
12
|
+
from .models import (
|
|
13
|
+
AssumeRoleNotSupported,
|
|
14
|
+
CloudIdentity,
|
|
15
|
+
CloudProviderNotFound,
|
|
16
|
+
CloudProviderType,
|
|
17
|
+
JWTType,
|
|
18
|
+
ProviderIdentityUnavailable,
|
|
19
|
+
ProviderNotDetected,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
__all__ = [
|
|
23
|
+
"detect_provider",
|
|
24
|
+
"DETECT_PROVIDER_DEFAULT_TIMEOUT",
|
|
25
|
+
"get_jwt",
|
|
26
|
+
"get_jwt_database",
|
|
27
|
+
"get_jwt_api",
|
|
28
|
+
"CloudIdentity",
|
|
29
|
+
"CloudProviderType",
|
|
30
|
+
"JWTType",
|
|
31
|
+
"CloudProviderNotFound",
|
|
32
|
+
"ProviderNotDetected",
|
|
33
|
+
"ProviderIdentityUnavailable",
|
|
34
|
+
"AssumeRoleNotSupported",
|
|
35
|
+
]
|
s2iam/api.py
ADDED
|
@@ -0,0 +1,300 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Main API for the s2iam library.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import os
|
|
7
|
+
import queue
|
|
8
|
+
import threading
|
|
9
|
+
import time
|
|
10
|
+
from typing import Any, Dict, List, NoReturn, Optional
|
|
11
|
+
|
|
12
|
+
from .aws import new_client as new_aws_client
|
|
13
|
+
from .azure import new_client as new_azure_client
|
|
14
|
+
from .gcp import new_client as new_gcp_client
|
|
15
|
+
from .models import CloudProviderClient, CloudProviderNotFound, Logger
|
|
16
|
+
|
|
17
|
+
DETECT_PROVIDER_DEFAULT_TIMEOUT: float = 10.0
|
|
18
|
+
"""Default timeout (seconds) for provider detection.
|
|
19
|
+
|
|
20
|
+
Rationale: Prefer avoiding false negatives over minimizing worst‑case wait.
|
|
21
|
+
10s ceiling handles slow or throttled metadata on constrained CI VMs while
|
|
22
|
+
early success still returns in sub‑second typical cases. Mirrors Go/Java parity.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class DefaultLogger:
|
|
27
|
+
"""Default logger implementation."""
|
|
28
|
+
|
|
29
|
+
def log(self, message: str) -> None:
|
|
30
|
+
"""Log a message to stdout."""
|
|
31
|
+
print(f"[s2iam] {message}")
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
async def detect_provider(
|
|
35
|
+
timeout: float = DETECT_PROVIDER_DEFAULT_TIMEOUT,
|
|
36
|
+
logger: Optional[Logger] = None,
|
|
37
|
+
clients: Optional[list[CloudProviderClient]] = None,
|
|
38
|
+
) -> CloudProviderClient:
|
|
39
|
+
"""
|
|
40
|
+
Detect which cloud provider we're running on.
|
|
41
|
+
|
|
42
|
+
Args:
|
|
43
|
+
timeout: Detection timeout in seconds
|
|
44
|
+
logger: Optional logger instance
|
|
45
|
+
clients: Optional list of custom provider clients
|
|
46
|
+
|
|
47
|
+
Returns:
|
|
48
|
+
CloudProviderClient for the detected provider
|
|
49
|
+
|
|
50
|
+
Raises:
|
|
51
|
+
CloudProviderNotFound: If no provider can be detected
|
|
52
|
+
"""
|
|
53
|
+
# Set up logger only if explicit debugging flag is set; production code must not branch
|
|
54
|
+
# on test harness-only environment variables. Rich diagnostics are instead
|
|
55
|
+
# surfaced via aggregated exception messages below.
|
|
56
|
+
debugging = os.environ.get("S2IAM_DEBUGGING", "").lower() == "true"
|
|
57
|
+
debug_timing = os.environ.get("S2IAM_DEBUG_TIMING", "").lower() == "true"
|
|
58
|
+
if logger is None and debugging:
|
|
59
|
+
logger = DefaultLogger()
|
|
60
|
+
|
|
61
|
+
# Create default clients if none provided
|
|
62
|
+
if clients is None:
|
|
63
|
+
clients = [
|
|
64
|
+
new_aws_client(logger),
|
|
65
|
+
new_gcp_client(logger),
|
|
66
|
+
new_azure_client(logger),
|
|
67
|
+
]
|
|
68
|
+
|
|
69
|
+
# Phase 1: fast_detect sequentially (purely local; must not do network I/O).
|
|
70
|
+
for c in clients:
|
|
71
|
+
try:
|
|
72
|
+
await c.fast_detect()
|
|
73
|
+
if logger:
|
|
74
|
+
logger.log(f"Fast detected provider: {c.get_type().value}")
|
|
75
|
+
return c
|
|
76
|
+
except Exception:
|
|
77
|
+
# Not detected via fast path; move to next provider.
|
|
78
|
+
continue
|
|
79
|
+
|
|
80
|
+
# Phase 2: full detection using threads (mirrors Go goroutines + first-winner channel).
|
|
81
|
+
# Invariants relied upon here: each client's detect() MUST raise on negative outcome; only a
|
|
82
|
+
# positively detected client (internal flag set) returns normally. This prevents selecting a
|
|
83
|
+
# provider that will later fail with ProviderNotDetected when building identity headers.
|
|
84
|
+
result_queue: "queue.Queue[CloudProviderClient]" = queue.Queue()
|
|
85
|
+
stop_event = threading.Event()
|
|
86
|
+
all_errors: list[str] = []
|
|
87
|
+
errors_lock = threading.Lock()
|
|
88
|
+
# Structured per-provider status for enhanced error reporting.
|
|
89
|
+
# Each element: {provider, status=success|error|timeout|skipped, elapsed_ms?, error?}
|
|
90
|
+
provider_status: List[Dict[str, Any]] = []
|
|
91
|
+
status_lock = threading.Lock()
|
|
92
|
+
|
|
93
|
+
def record_status(entry: Dict[str, Any]) -> None:
|
|
94
|
+
with status_lock:
|
|
95
|
+
provider_status.append(entry)
|
|
96
|
+
|
|
97
|
+
# Track per-thread event loops so we can cancel/stop them on global timeout
|
|
98
|
+
provider_loops: list[asyncio.AbstractEventLoop] = []
|
|
99
|
+
loops_lock = threading.Lock()
|
|
100
|
+
|
|
101
|
+
def test_provider_sync(client: CloudProviderClient) -> None:
|
|
102
|
+
"""Test a provider in a thread (like Go goroutine)."""
|
|
103
|
+
if stop_event.is_set():
|
|
104
|
+
record_status({"provider": client.get_type().value, "status": "skipped"})
|
|
105
|
+
return
|
|
106
|
+
thread_start = time.monotonic()
|
|
107
|
+
if logger and (debugging or debug_timing):
|
|
108
|
+
logger.log(f"DETECT_THREAD_START provider={client.get_type().value} outer_timeout_s={timeout}")
|
|
109
|
+
try:
|
|
110
|
+
# Run the async detect() in this thread's event loop
|
|
111
|
+
loop = asyncio.new_event_loop()
|
|
112
|
+
asyncio.set_event_loop(loop)
|
|
113
|
+
with loops_lock:
|
|
114
|
+
provider_loops.append(loop)
|
|
115
|
+
try:
|
|
116
|
+
detect_coro = client.detect()
|
|
117
|
+
loop.run_until_complete(detect_coro)
|
|
118
|
+
# Post-call validation: provider must internally mark detected.
|
|
119
|
+
is_detected = True
|
|
120
|
+
try:
|
|
121
|
+
# Allow provider to expose explicit method; fallback to attribute.
|
|
122
|
+
if hasattr(client, "is_detected") and callable(getattr(client, "is_detected")):
|
|
123
|
+
is_detected = bool(client.is_detected())
|
|
124
|
+
elif hasattr(client, "_detected"):
|
|
125
|
+
is_detected = bool(getattr(client, "_detected"))
|
|
126
|
+
except Exception: # noqa: BLE001 - conservative: treat as success only if flag readable
|
|
127
|
+
pass
|
|
128
|
+
elapsed_ms = int((time.monotonic() - thread_start) * 1000)
|
|
129
|
+
if is_detected and not stop_event.is_set():
|
|
130
|
+
result_queue.put(client)
|
|
131
|
+
stop_event.set()
|
|
132
|
+
record_status(
|
|
133
|
+
{
|
|
134
|
+
"provider": client.get_type().value,
|
|
135
|
+
"status": "success",
|
|
136
|
+
"elapsed_ms": elapsed_ms,
|
|
137
|
+
}
|
|
138
|
+
)
|
|
139
|
+
if logger and (debugging or debug_timing):
|
|
140
|
+
logger.log(f"DETECT_THREAD_SUCCESS provider={client.get_type().value} elapsed_ms={elapsed_ms}")
|
|
141
|
+
else:
|
|
142
|
+
record_status(
|
|
143
|
+
{
|
|
144
|
+
"provider": client.get_type().value,
|
|
145
|
+
"status": "not_detected",
|
|
146
|
+
"elapsed_ms": elapsed_ms,
|
|
147
|
+
}
|
|
148
|
+
)
|
|
149
|
+
if logger and (debugging or debug_timing):
|
|
150
|
+
logger.log(
|
|
151
|
+
f"DETECT_THREAD_NOT_DETECTED provider={client.get_type().value} elapsed_ms={elapsed_ms}"
|
|
152
|
+
)
|
|
153
|
+
finally:
|
|
154
|
+
loop.close()
|
|
155
|
+
except Exception as e:
|
|
156
|
+
with errors_lock:
|
|
157
|
+
all_errors.append(f"Provider {client.get_type().value} detection failed: {e}")
|
|
158
|
+
elapsed_ms = int((time.monotonic() - thread_start) * 1000)
|
|
159
|
+
record_status(
|
|
160
|
+
{
|
|
161
|
+
"provider": client.get_type().value,
|
|
162
|
+
"status": "error",
|
|
163
|
+
"elapsed_ms": elapsed_ms,
|
|
164
|
+
"error": str(e)[:400],
|
|
165
|
+
}
|
|
166
|
+
)
|
|
167
|
+
if logger and (debugging or debug_timing):
|
|
168
|
+
logger.log(f"DETECT_THREAD_ERROR provider={client.get_type().value} elapsed_ms={elapsed_ms} error={e}")
|
|
169
|
+
|
|
170
|
+
# Start threads for each provider (like Go goroutines)
|
|
171
|
+
threads = []
|
|
172
|
+
for client in clients:
|
|
173
|
+
thread = threading.Thread(target=test_provider_sync, args=(client,))
|
|
174
|
+
thread.daemon = True
|
|
175
|
+
thread.start()
|
|
176
|
+
threads.append(thread)
|
|
177
|
+
|
|
178
|
+
# Wait for first result or timeout (like Go select)
|
|
179
|
+
detection_start = time.monotonic()
|
|
180
|
+
# Track how many threads have finished (success or error) to allow early exit when all done.
|
|
181
|
+
total_clients = len(clients)
|
|
182
|
+
# Poll loop instead of single blocking get so we can detect early-failure condition.
|
|
183
|
+
remaining = timeout
|
|
184
|
+
interval = 0.05 # 50ms poll granularity (balance: coarse enough to reduce wakeups, fine enough for fast success)
|
|
185
|
+
while remaining > 0:
|
|
186
|
+
start_poll = time.monotonic()
|
|
187
|
+
try:
|
|
188
|
+
result: CloudProviderClient = result_queue.get(timeout=min(interval, remaining))
|
|
189
|
+
stop_event.set() # Ensure all threads stop
|
|
190
|
+
total_elapsed_ms = int((time.monotonic() - detection_start) * 1000)
|
|
191
|
+
if logger:
|
|
192
|
+
if debugging or debug_timing:
|
|
193
|
+
logger.log(
|
|
194
|
+
(
|
|
195
|
+
"DETECT_COMPLETE status=success "
|
|
196
|
+
f"provider={result.get_type().value} "
|
|
197
|
+
f"total_elapsed_ms={total_elapsed_ms}"
|
|
198
|
+
)
|
|
199
|
+
)
|
|
200
|
+
else:
|
|
201
|
+
logger.log(f"Detected provider: {result.get_type().value}")
|
|
202
|
+
return result
|
|
203
|
+
except queue.Empty:
|
|
204
|
+
pass
|
|
205
|
+
# Early failure: if every thread has produced a terminal status (success/error/timeout/skipped)
|
|
206
|
+
with status_lock:
|
|
207
|
+
finished = sum(1 for ps in provider_status if ps["status"] in {"success", "error", "timeout", "skipped"})
|
|
208
|
+
if finished >= total_clients and not any(ps["status"] == "success" for ps in provider_status):
|
|
209
|
+
# All threads ended without success -> raise immediately (no need to wait remaining timeout)
|
|
210
|
+
break
|
|
211
|
+
remaining -= time.monotonic() - start_poll
|
|
212
|
+
else:
|
|
213
|
+
# Loop ended naturally (remaining <= 0) without a success result
|
|
214
|
+
if logger and (debugging or debug_timing):
|
|
215
|
+
logger.log(
|
|
216
|
+
"DETECT_LOOP_COMPLETE no-success reason=timeout-before-result "
|
|
217
|
+
f"elapsed_ms={int((time.monotonic()-detection_start)*1000)}"
|
|
218
|
+
)
|
|
219
|
+
|
|
220
|
+
# No provider detected within timeout or all failed fast.
|
|
221
|
+
stop_event.set() # Ensure all threads stop
|
|
222
|
+
|
|
223
|
+
# If we broke out of loop without returning (early failure or timeout), raise composed error.
|
|
224
|
+
_raise_detection_timeout(
|
|
225
|
+
timeout=timeout,
|
|
226
|
+
detection_start=detection_start,
|
|
227
|
+
clients=clients,
|
|
228
|
+
threads=threads,
|
|
229
|
+
provider_status=provider_status,
|
|
230
|
+
status_lock=status_lock,
|
|
231
|
+
all_errors=all_errors,
|
|
232
|
+
logger=logger,
|
|
233
|
+
debugging=debugging,
|
|
234
|
+
debug_timing=debug_timing,
|
|
235
|
+
stop_event=stop_event,
|
|
236
|
+
provider_loops=provider_loops,
|
|
237
|
+
loops_lock=loops_lock,
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def _raise_detection_timeout(
|
|
242
|
+
*,
|
|
243
|
+
timeout: float,
|
|
244
|
+
detection_start: float,
|
|
245
|
+
clients: list[CloudProviderClient],
|
|
246
|
+
threads: list[threading.Thread],
|
|
247
|
+
provider_status: list[dict[str, Any]],
|
|
248
|
+
status_lock: threading.Lock,
|
|
249
|
+
all_errors: list[str],
|
|
250
|
+
logger: Optional[Logger],
|
|
251
|
+
debugging: bool,
|
|
252
|
+
debug_timing: bool,
|
|
253
|
+
stop_event: threading.Event,
|
|
254
|
+
provider_loops: list[asyncio.AbstractEventLoop],
|
|
255
|
+
loops_lock: threading.Lock,
|
|
256
|
+
) -> NoReturn:
|
|
257
|
+
"""Compose and raise CloudProviderNotFound for a detection timeout.
|
|
258
|
+
|
|
259
|
+
Isolated to keep the main detect_provider flow skimmable and ease future
|
|
260
|
+
experiments (e.g., per-provider granular timeouts or retry policy integration).
|
|
261
|
+
"""
|
|
262
|
+
stop_event.set()
|
|
263
|
+
# Attempt to stop any active provider event loops to prevent post-timeout drift
|
|
264
|
+
with loops_lock:
|
|
265
|
+
for loop in provider_loops:
|
|
266
|
+
if loop.is_running():
|
|
267
|
+
try:
|
|
268
|
+
loop.call_soon_threadsafe(loop.stop)
|
|
269
|
+
except Exception: # noqa: BLE001 - best effort cancellation
|
|
270
|
+
pass
|
|
271
|
+
for thread in threads:
|
|
272
|
+
# 50ms join chosen: keeps worst-case added delay bounded (< provider timeout granularity)
|
|
273
|
+
# while giving threads a chance to observe stop_event and exit cleanly.
|
|
274
|
+
thread.join(timeout=0.05)
|
|
275
|
+
total_elapsed_ms = int((time.monotonic() - detection_start) * 1000)
|
|
276
|
+
with status_lock:
|
|
277
|
+
known = {p["provider"] for p in provider_status}
|
|
278
|
+
for c in clients:
|
|
279
|
+
name = c.get_type().value
|
|
280
|
+
if name not in known:
|
|
281
|
+
provider_status.append({"provider": name, "status": "timeout"})
|
|
282
|
+
if logger and (debugging or debug_timing):
|
|
283
|
+
joined_errors_dbg = " | ".join(all_errors)[:800]
|
|
284
|
+
logger.log(
|
|
285
|
+
(
|
|
286
|
+
"DETECT_COMPLETE status=timeout "
|
|
287
|
+
f"total_elapsed_ms={total_elapsed_ms} timeout_s={timeout} "
|
|
288
|
+
f"errors='{joined_errors_dbg}'"
|
|
289
|
+
)
|
|
290
|
+
)
|
|
291
|
+
summary = ", ".join(
|
|
292
|
+
f"{ps['provider']}:{ps['status']}{('@'+str(ps['elapsed_ms'])+'ms') if 'elapsed_ms' in ps else ''}"
|
|
293
|
+
for ps in provider_status
|
|
294
|
+
)
|
|
295
|
+
errors_str = " | ".join(all_errors)[:800] if all_errors else "<no-provider-errors>"
|
|
296
|
+
raise CloudProviderNotFound(
|
|
297
|
+
"Provider detection timed out: "
|
|
298
|
+
f"timeout_s={timeout} total_elapsed_ms={total_elapsed_ms} providers={len(clients)} "
|
|
299
|
+
f"error_count={len(all_errors)} provider_status=[{summary}] errors=[{errors_str}]"
|
|
300
|
+
)
|
s2iam/aws/__init__.py
ADDED
|
@@ -0,0 +1,271 @@
|
|
|
1
|
+
"""AWS cloud provider client implementation.
|
|
2
|
+
|
|
3
|
+
Single implementation aligned with the Go reference: fast env/IMDS detection,
|
|
4
|
+
STS fallback, optional role assumption, region derivation, and identity header
|
|
5
|
+
generation.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import asyncio
|
|
9
|
+
import os
|
|
10
|
+
from typing import Any, Optional
|
|
11
|
+
|
|
12
|
+
from ..models import (
|
|
13
|
+
CloudIdentity,
|
|
14
|
+
CloudProviderClient,
|
|
15
|
+
CloudProviderType,
|
|
16
|
+
Logger,
|
|
17
|
+
ProviderIdentityUnavailable,
|
|
18
|
+
ProviderNotDetected,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class AWSClient(CloudProviderClient):
|
|
23
|
+
_logger: Optional[Logger]
|
|
24
|
+
_detected: bool
|
|
25
|
+
_region: Optional[str]
|
|
26
|
+
_identity: Optional[CloudIdentity]
|
|
27
|
+
_role_arn: Optional[str]
|
|
28
|
+
_sts_client: Optional[Any]
|
|
29
|
+
_session: Optional[Any]
|
|
30
|
+
|
|
31
|
+
def __init__(self, logger: Optional[Logger] = None):
|
|
32
|
+
self._logger = logger
|
|
33
|
+
self._detected = False
|
|
34
|
+
self._region = None
|
|
35
|
+
self._identity = None
|
|
36
|
+
self._role_arn = None
|
|
37
|
+
self._sts_client = None
|
|
38
|
+
self._session = None
|
|
39
|
+
|
|
40
|
+
def _log(self, message: str) -> None:
|
|
41
|
+
if self._logger:
|
|
42
|
+
self._logger.log(f"AWS: {message}")
|
|
43
|
+
|
|
44
|
+
async def _check_metadata_service(self) -> bool:
|
|
45
|
+
"""Best-effort IMDSv2 then IMDSv1 probe (<= ~3s worst case)."""
|
|
46
|
+
try: # noqa: BLE001
|
|
47
|
+
import aiohttp
|
|
48
|
+
|
|
49
|
+
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=3)) as session:
|
|
50
|
+
async with session.put(
|
|
51
|
+
"http://169.254.169.254/latest/api/token",
|
|
52
|
+
headers={"X-aws-ec2-metadata-token-ttl-seconds": "21600"},
|
|
53
|
+
) as token_resp:
|
|
54
|
+
if token_resp.status == 200:
|
|
55
|
+
token = await token_resp.text()
|
|
56
|
+
async with session.get(
|
|
57
|
+
"http://169.254.169.254/latest/meta-data/instance-id",
|
|
58
|
+
headers={"X-aws-ec2-metadata-token": token},
|
|
59
|
+
) as resp:
|
|
60
|
+
if resp.status == 200:
|
|
61
|
+
return True
|
|
62
|
+
|
|
63
|
+
async with session.get(
|
|
64
|
+
"http://169.254.169.254/latest/meta-data/instance-id",
|
|
65
|
+
timeout=aiohttp.ClientTimeout(total=2),
|
|
66
|
+
) as resp:
|
|
67
|
+
return resp.status == 200
|
|
68
|
+
except Exception as e: # noqa: BLE001
|
|
69
|
+
self._log(f"Metadata service check failed: {e}")
|
|
70
|
+
return False
|
|
71
|
+
|
|
72
|
+
async def detect(self) -> None:
|
|
73
|
+
# Full (network-inclusive) detection. Raise on failure so orchestrator never
|
|
74
|
+
# selects an undetected client (prevents later ProviderNotDetected errors).
|
|
75
|
+
self._log("Starting AWS detection (full phase)")
|
|
76
|
+
if self._detected:
|
|
77
|
+
return
|
|
78
|
+
|
|
79
|
+
if await self._check_metadata_service():
|
|
80
|
+
self._detected = True
|
|
81
|
+
self._log("Detected via metadata service")
|
|
82
|
+
return
|
|
83
|
+
|
|
84
|
+
try:
|
|
85
|
+
import boto3 # optional dependency in some usage contexts
|
|
86
|
+
except ImportError as e: # noqa: BLE001
|
|
87
|
+
self._log(f"boto3 import failed: {e}")
|
|
88
|
+
else:
|
|
89
|
+
try: # noqa: BLE001
|
|
90
|
+
sts_client = boto3.client("sts")
|
|
91
|
+
identity = sts_client.get_caller_identity()
|
|
92
|
+
if identity.get("Account"):
|
|
93
|
+
self._detected = True
|
|
94
|
+
self._log("Detected via STS")
|
|
95
|
+
return
|
|
96
|
+
except Exception as e: # noqa: BLE001
|
|
97
|
+
self._log(f"STS detection failed: {e}")
|
|
98
|
+
|
|
99
|
+
self._log("AWS full detection did not succeed; raising")
|
|
100
|
+
raise Exception("AWS provider not detected")
|
|
101
|
+
|
|
102
|
+
async def fast_detect(self) -> None:
|
|
103
|
+
"""Fast detection: env only, no network calls."""
|
|
104
|
+
# IRSA / web identity short-circuit: ONLY honor explicit env vars.
|
|
105
|
+
if os.environ.get("AWS_WEB_IDENTITY_TOKEN_FILE") or os.environ.get("AWS_ROLE_ARN"):
|
|
106
|
+
self._detected = True
|
|
107
|
+
self._log("FastDetect: IRSA environment variables present")
|
|
108
|
+
return
|
|
109
|
+
|
|
110
|
+
for var in ("AWS_EXECUTION_ENV", "AWS_REGION", "AWS_DEFAULT_REGION", "AWS_LAMBDA_FUNCTION_NAME"):
|
|
111
|
+
if os.environ.get(var):
|
|
112
|
+
self._detected = True
|
|
113
|
+
self._log(f"FastDetect: detected via env var {var}")
|
|
114
|
+
return
|
|
115
|
+
raise Exception("FastDetect: no AWS indicators")
|
|
116
|
+
|
|
117
|
+
async def _ensure_region(self) -> None:
|
|
118
|
+
if self._region:
|
|
119
|
+
return
|
|
120
|
+
|
|
121
|
+
region = os.environ.get("AWS_REGION") or os.environ.get("AWS_DEFAULT_REGION")
|
|
122
|
+
if not region:
|
|
123
|
+
try: # noqa: BLE001
|
|
124
|
+
import aiohttp
|
|
125
|
+
|
|
126
|
+
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=3)) as session:
|
|
127
|
+
async with session.put(
|
|
128
|
+
"http://169.254.169.254/latest/api/token",
|
|
129
|
+
headers={"X-aws-ec2-metadata-token-ttl-seconds": "21600"},
|
|
130
|
+
) as token_resp:
|
|
131
|
+
if token_resp.status == 200:
|
|
132
|
+
token = await token_resp.text()
|
|
133
|
+
async with session.get(
|
|
134
|
+
"http://169.254.169.254/latest/meta-data/placement/region",
|
|
135
|
+
headers={"X-aws-ec2-metadata-token": token},
|
|
136
|
+
) as region_resp:
|
|
137
|
+
if region_resp.status == 200:
|
|
138
|
+
region = await region_resp.text()
|
|
139
|
+
self._log(f"Region from metadata: {region}")
|
|
140
|
+
except Exception as e: # noqa: BLE001
|
|
141
|
+
self._log(f"Region metadata lookup failed: {e}")
|
|
142
|
+
|
|
143
|
+
if not region:
|
|
144
|
+
region = "us-east-1"
|
|
145
|
+
self._log("Defaulting region to us-east-1")
|
|
146
|
+
self._region = region
|
|
147
|
+
|
|
148
|
+
def get_type(self) -> CloudProviderType:
|
|
149
|
+
return CloudProviderType.AWS
|
|
150
|
+
|
|
151
|
+
def assume_role(self, role_identifier: str) -> "AWSClient":
|
|
152
|
+
clone = AWSClient(self._logger)
|
|
153
|
+
clone._detected = self._detected
|
|
154
|
+
clone._region = self._region
|
|
155
|
+
clone._role_arn = role_identifier
|
|
156
|
+
clone._sts_client = self._sts_client
|
|
157
|
+
return clone
|
|
158
|
+
|
|
159
|
+
async def get_identity_headers(
|
|
160
|
+
self, additional_params: Optional[dict[str, str]] = None
|
|
161
|
+
) -> tuple[dict[str, str], CloudIdentity]: # noqa: D401,E501
|
|
162
|
+
if not self._detected:
|
|
163
|
+
raise ProviderNotDetected("AWS provider not detected, call detect() first")
|
|
164
|
+
|
|
165
|
+
if not self._sts_client:
|
|
166
|
+
import boto3
|
|
167
|
+
|
|
168
|
+
await self._ensure_region()
|
|
169
|
+
self._session = boto3.Session()
|
|
170
|
+
self._sts_client = self._session.client("sts", region_name=self._region)
|
|
171
|
+
self._log("Initialized STS client")
|
|
172
|
+
|
|
173
|
+
if self._sts_client is None or self._session is None:
|
|
174
|
+
# Can happen when clone created via assume_role (sts_client copied but session not)
|
|
175
|
+
import boto3
|
|
176
|
+
|
|
177
|
+
if self._session is None:
|
|
178
|
+
self._session = boto3.Session()
|
|
179
|
+
if self._sts_client is None:
|
|
180
|
+
self._sts_client = self._session.client("sts", region_name=self._region)
|
|
181
|
+
self._log("Recovered missing STS session/client state")
|
|
182
|
+
|
|
183
|
+
# Narrow optionals after explicit check
|
|
184
|
+
sts_client = self._sts_client
|
|
185
|
+
session_obj = self._session
|
|
186
|
+
|
|
187
|
+
loop = asyncio.get_event_loop()
|
|
188
|
+
|
|
189
|
+
try: # noqa: BLE001
|
|
190
|
+
if self._role_arn:
|
|
191
|
+
self._log(f"Assuming role {self._role_arn}")
|
|
192
|
+
assume_resp = await loop.run_in_executor(
|
|
193
|
+
None,
|
|
194
|
+
lambda: sts_client.assume_role(
|
|
195
|
+
RoleArn=self._role_arn,
|
|
196
|
+
RoleSessionName="s2iam-session",
|
|
197
|
+
),
|
|
198
|
+
)
|
|
199
|
+
creds = assume_resp["Credentials"]
|
|
200
|
+
import boto3
|
|
201
|
+
|
|
202
|
+
assumed_session = boto3.Session(
|
|
203
|
+
aws_access_key_id=creds["AccessKeyId"],
|
|
204
|
+
aws_secret_access_key=creds["SecretAccessKey"],
|
|
205
|
+
aws_session_token=creds["SessionToken"],
|
|
206
|
+
region_name=self._region,
|
|
207
|
+
)
|
|
208
|
+
assumed_sts = assumed_session.client("sts")
|
|
209
|
+
identity_resp = await loop.run_in_executor(None, assumed_sts.get_caller_identity)
|
|
210
|
+
headers = {
|
|
211
|
+
"X-AWS-Access-Key-ID": creds["AccessKeyId"],
|
|
212
|
+
"X-AWS-Secret-Access-Key": creds["SecretAccessKey"],
|
|
213
|
+
"X-AWS-Session-Token": creds["SessionToken"],
|
|
214
|
+
}
|
|
215
|
+
else:
|
|
216
|
+
identity_resp = await loop.run_in_executor(None, sts_client.get_caller_identity)
|
|
217
|
+
role_assumed = (
|
|
218
|
+
":assumed-role/" in identity_resp["Arn"] or os.environ.get("AWS_SESSION_TOKEN") is not None
|
|
219
|
+
)
|
|
220
|
+
if role_assumed:
|
|
221
|
+
creds = session_obj.get_credentials()
|
|
222
|
+
headers = {
|
|
223
|
+
"X-AWS-Access-Key-ID": creds.access_key,
|
|
224
|
+
"X-AWS-Secret-Access-Key": creds.secret_key,
|
|
225
|
+
"X-Cloud-Provider": "aws",
|
|
226
|
+
}
|
|
227
|
+
if creds.token:
|
|
228
|
+
headers["X-AWS-Session-Token"] = creds.token
|
|
229
|
+
if os.environ.get("AWS_WEB_IDENTITY_TOKEN_FILE") or os.environ.get("AWS_ROLE_ARN"):
|
|
230
|
+
self._log("Using IRSA web identity session credentials")
|
|
231
|
+
else:
|
|
232
|
+
self._log("Getting session token for static credentials")
|
|
233
|
+
session_resp = await loop.run_in_executor(None, sts_client.get_session_token)
|
|
234
|
+
sc = session_resp["Credentials"]
|
|
235
|
+
headers = {
|
|
236
|
+
"X-AWS-Access-Key-ID": sc["AccessKeyId"],
|
|
237
|
+
"X-AWS-Secret-Access-Key": sc["SecretAccessKey"],
|
|
238
|
+
"X-AWS-Session-Token": sc["SessionToken"],
|
|
239
|
+
"X-Cloud-Provider": "aws",
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
arn = identity_resp["Arn"]
|
|
243
|
+
parts = arn.split(":")
|
|
244
|
+
region_from_arn = parts[3] if len(parts) > 3 else ""
|
|
245
|
+
resource_type = ""
|
|
246
|
+
if len(parts) > 5:
|
|
247
|
+
res_parts = parts[5].split("/")
|
|
248
|
+
if res_parts and res_parts[0]:
|
|
249
|
+
resource_type = res_parts[0]
|
|
250
|
+
|
|
251
|
+
# If region unset locally (IRSA path without env/metadata), adopt ARN region
|
|
252
|
+
if not self._region and region_from_arn:
|
|
253
|
+
self._region = region_from_arn
|
|
254
|
+
self._log(f"Derived region from ARN: {self._region}")
|
|
255
|
+
|
|
256
|
+
identity = CloudIdentity(
|
|
257
|
+
provider=CloudProviderType.AWS,
|
|
258
|
+
identifier=arn,
|
|
259
|
+
account_id=identity_resp["Account"],
|
|
260
|
+
region=region_from_arn,
|
|
261
|
+
resource_type=resource_type,
|
|
262
|
+
)
|
|
263
|
+
self._log(f"Generated headers for identity: {identity.identifier}")
|
|
264
|
+
return headers, identity
|
|
265
|
+
except Exception as e: # noqa: BLE001
|
|
266
|
+
self._log(f"Failed to build identity headers: {e}")
|
|
267
|
+
raise ProviderIdentityUnavailable(f"Failed to get AWS identity: {e}")
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
def new_client(logger: Optional[Logger] = None) -> CloudProviderClient:
|
|
271
|
+
return AWSClient(logger)
|