mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test: real-cloud EC2Provider tests — slow-marked, gated by BYOC env vars
Validations #3 (real boot), #11 (real BYOC fail-fast), and the real-cloud piece of #8. Skipped by default; enable with `pytest -m slow` when the `LITELLM_AGENT_AWS_*` + `LITELLM_TEST_*` env vars are set. Each test wraps RunInstances in try/finally with TerminateInstances and installs a 60-min process watchdog (per the AWS safety boundary) so a hung test cannot leak an instance. LIT-2878
This commit is contained in:
parent
0b3cb3040d
commit
48c8ba1fe7
1 changed files with 208 additions and 0 deletions
|
|
@ -0,0 +1,208 @@
|
|||
"""
|
||||
Real-cloud tests for `EC2Provider`.
|
||||
|
||||
Skipped by default. Enable with `pytest -m slow` and the BYOC env vars set:
|
||||
|
||||
LITELLM_AGENT_AWS_ACCESS_KEY_ID
|
||||
LITELLM_AGENT_AWS_SECRET_ACCESS_KEY
|
||||
LITELLM_AGENT_AWS_REGION (default us-west-2)
|
||||
LITELLM_TEST_SUBNET_ID
|
||||
LITELLM_TEST_SECURITY_GROUP_ID
|
||||
LITELLM_TEST_IAM_INSTANCE_PROFILE
|
||||
LITELLM_TEST_AMI_ID
|
||||
|
||||
These are the resources B0 captured for the BYOC PoC account
|
||||
(see LIT-2888 deliverables comment).
|
||||
|
||||
Each test is wrapped in a try/finally that calls TerminateInstances. A 60-min
|
||||
process-wide watchdog is also installed so a hung test cannot leak an
|
||||
instance overnight.
|
||||
|
||||
Cost per run: ~$0.03 per session (one t3.large for under a minute).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from typing import List, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.agent_session_endpoints.vm_providers import (
|
||||
AwsCreds,
|
||||
EC2Provider,
|
||||
Ec2Config,
|
||||
ProvisionContext,
|
||||
Repo,
|
||||
VMHandle,
|
||||
VMState,
|
||||
)
|
||||
|
||||
|
||||
# Test markers — `slow` is the standard real-cloud gate.
|
||||
pytestmark = [pytest.mark.slow]
|
||||
|
||||
|
||||
def _have_creds() -> bool:
|
||||
return bool(
|
||||
os.getenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID")
|
||||
and os.getenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY")
|
||||
and os.getenv("LITELLM_TEST_AMI_ID")
|
||||
)
|
||||
|
||||
|
||||
_skip_no_creds = pytest.mark.skipif(
|
||||
not _have_creds(),
|
||||
reason="real-cloud test requires LITELLM_AGENT_AWS_* + LITELLM_TEST_* env vars",
|
||||
)
|
||||
|
||||
|
||||
def _ec2_config_from_env() -> Ec2Config:
|
||||
return Ec2Config(
|
||||
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
|
||||
subnet_id=os.getenv("LITELLM_TEST_SUBNET_ID"),
|
||||
security_group_id=os.getenv("LITELLM_TEST_SECURITY_GROUP_ID"),
|
||||
iam_instance_profile=os.getenv("LITELLM_TEST_IAM_INSTANCE_PROFILE"),
|
||||
instance_type="t3.large",
|
||||
use_spot=False, # real tests prefer determinism over savings
|
||||
ami_id=os.getenv("LITELLM_TEST_AMI_ID"),
|
||||
)
|
||||
|
||||
|
||||
def _aws_creds_from_env() -> AwsCreds:
|
||||
return AwsCreds(
|
||||
access_key_id=os.environ["LITELLM_AGENT_AWS_ACCESS_KEY_ID"],
|
||||
secret_access_key=os.environ["LITELLM_AGENT_AWS_SECRET_ACCESS_KEY"],
|
||||
session_token=os.getenv("LITELLM_AGENT_AWS_SESSION_TOKEN"),
|
||||
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
|
||||
)
|
||||
|
||||
|
||||
def _install_watchdog(
|
||||
provider: EC2Provider, handles: List[VMHandle]
|
||||
) -> threading.Timer:
|
||||
"""60-min hard watchdog: terminate everything we tracked, no matter what."""
|
||||
|
||||
def _kill_all() -> None:
|
||||
import asyncio # noqa: PLC0415
|
||||
|
||||
async def _go():
|
||||
for h in handles:
|
||||
try:
|
||||
await provider.terminate(h, aws_creds=_aws_creds_from_env())
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
asyncio.run(_go())
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
timer = threading.Timer(60 * 60, _kill_all)
|
||||
timer.daemon = True
|
||||
timer.start()
|
||||
return timer
|
||||
|
||||
|
||||
@_skip_no_creds
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_boot():
|
||||
"""Validation #3: instance reaches `running` within 90s, then we terminate."""
|
||||
handles: List[VMHandle] = []
|
||||
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
|
||||
watchdog = _install_watchdog(provider, handles)
|
||||
|
||||
ctx = ProvisionContext(
|
||||
session_id=f"sess-test-real-{int(time.time())}",
|
||||
team_id="team-test-real",
|
||||
agent_id="agent-test-real",
|
||||
repos=[Repo(url="https://github.com/octocat/Hello-World")],
|
||||
env_vars={"FOO": "bar"},
|
||||
aws_creds=_aws_creds_from_env(),
|
||||
ec2_config=_ec2_config_from_env(),
|
||||
daemon_jwt="not-a-real-jwt-test-only",
|
||||
daemon_base_url="https://example.invalid/",
|
||||
mode="session",
|
||||
)
|
||||
|
||||
handle: Optional[VMHandle] = None
|
||||
try:
|
||||
handle = await provider.provision(ctx)
|
||||
handles.append(handle)
|
||||
assert handle.vm_id.startswith("i-")
|
||||
|
||||
# Poll for `running` (max 90s).
|
||||
deadline = time.time() + 90
|
||||
last: Optional[VMState] = None
|
||||
while time.time() < deadline:
|
||||
status = await provider.status(handle, aws_creds=ctx.aws_creds)
|
||||
last = status.state
|
||||
if status.state == VMState.RUNNING:
|
||||
break
|
||||
time.sleep(2)
|
||||
assert last == VMState.RUNNING, f"never reached running, last={last}"
|
||||
finally:
|
||||
if handle is not None:
|
||||
await provider.terminate(handle, aws_creds=ctx.aws_creds)
|
||||
watchdog.cancel()
|
||||
|
||||
|
||||
@_skip_no_creds
|
||||
@pytest.mark.asyncio
|
||||
async def test_byoc_invalid_creds_fail_fast():
|
||||
"""Validation #11: bad creds → InvalidCredentialsError within ~5s, no instance launched."""
|
||||
from litellm.proxy.agent_session_endpoints.vm_providers import (
|
||||
InvalidCredentialsError,
|
||||
)
|
||||
|
||||
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
|
||||
bad_creds = AwsCreds(
|
||||
access_key_id="AKIAINVALIDFAKEKEY00",
|
||||
secret_access_key="invalid-secret-do-not-leak",
|
||||
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
|
||||
)
|
||||
ctx = ProvisionContext(
|
||||
session_id=f"sess-bad-{int(time.time())}",
|
||||
team_id="team-bad",
|
||||
aws_creds=bad_creds,
|
||||
ec2_config=_ec2_config_from_env(),
|
||||
)
|
||||
t0 = time.time()
|
||||
with pytest.raises(InvalidCredentialsError):
|
||||
await provider.provision(ctx)
|
||||
elapsed = time.time() - t0
|
||||
assert elapsed < 10, f"fail-fast took too long: {elapsed:.1f}s"
|
||||
|
||||
|
||||
@_skip_no_creds
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_idempotent_real():
|
||||
"""Validation #8 piece: terminate the same instance twice — second is a no-op."""
|
||||
handles: List[VMHandle] = []
|
||||
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
|
||||
watchdog = _install_watchdog(provider, handles)
|
||||
|
||||
ctx = ProvisionContext(
|
||||
session_id=f"sess-term-{int(time.time())}",
|
||||
team_id="team-term",
|
||||
aws_creds=_aws_creds_from_env(),
|
||||
ec2_config=_ec2_config_from_env(),
|
||||
daemon_jwt="x",
|
||||
daemon_base_url="https://example.invalid/",
|
||||
)
|
||||
handle: Optional[VMHandle] = None
|
||||
try:
|
||||
handle = await provider.provision(ctx)
|
||||
handles.append(handle)
|
||||
await provider.terminate(handle, aws_creds=ctx.aws_creds)
|
||||
# Second call must not raise.
|
||||
await provider.terminate(handle, aws_creds=ctx.aws_creds)
|
||||
finally:
|
||||
if handle is not None:
|
||||
try:
|
||||
await provider.terminate(handle, aws_creds=ctx.aws_creds)
|
||||
except Exception:
|
||||
pass
|
||||
watchdog.cancel()
|
||||
Loading…
Add table
Reference in a new issue