Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 2 additions & 12 deletions src/tetra_rp/client.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import logging
from functools import wraps
from typing import List, Optional
from .core.resources import ServerlessResource, ResourceManager, NetworkVolume
from typing import List
from .core.resources import ServerlessResource, ResourceManager
from .stubs import stub_resource


Expand All @@ -12,7 +12,6 @@ def remote(
resource_config: ServerlessResource,
dependencies: List[str] = None,
system_dependencies: List[str] = None,
mount_volume: Optional[NetworkVolume] = None,
**extra,
):
"""
Expand Down Expand Up @@ -49,15 +48,6 @@ async def my_function(data):
def decorator(func):
@wraps(func)
async def wrapper(*args, **kwargs):
# Create netowrk volume if mount_volume is provided
if mount_volume:
try:
network_volume = await mount_volume.deploy()
resource_config.networkVolumeId = network_volume.id
except Exception as e:
log.error(f"Failed to create or mount network volume: {e}")
raise

resource_manager = ResourceManager()
remote_resource = await resource_manager.get_or_deploy_resource(
resource_config
Expand Down
36 changes: 9 additions & 27 deletions src/tetra_rp/core/api/runpod.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,12 @@
Bypasses the outdated runpod-python SDK limitations.
"""

import os
import json
import aiohttp
from typing import Dict, Any, Optional
import logging
import os
from typing import Any, Dict, Optional

import aiohttp

log = logging.getLogger(__name__)

Expand Down Expand Up @@ -267,31 +268,12 @@ async def _execute_rest(
raise Exception(f"HTTP request failed: {e}")

async def create_network_volume(self, payload: Dict[str, Any]) -> Dict[str, Any]:
"""
Create a network volume in Runpod.
"""Create a network volume in Runpod."""
log.debug(f"Creating network volume: {payload.get('name', 'unnamed')}")

Args:
datacenter_id (str): The ID of the datacenter where the volume will be created.
name (str): The name of the network volume.
size_gb (int): The size of the volume in GB.

Returns:
Dict[str, Any]: The created network volume details.
"""
datacenter_id = payload.get("dataCenterId")
if hasattr(datacenter_id, "value"):
# If datacenter_id is an enum, get its value
datacenter_id = datacenter_id.value
data = {
"dataCenterId": datacenter_id,
"name": payload.get("name"),
"size": payload.get("size"),
}
url = f"{RUNPOD_REST_API_URL}/networkvolumes"

log.debug(f"Creating network volume: {data.get('name', 'unnamed')}")

result = await self._execute_rest("POST", url, data)
result = await self._execute_rest(
"POST", f"{RUNPOD_REST_API_URL}/networkvolumes", payload
)

log.info(
f"Created network volume: {result.get('id', 'unknown')} - {result.get('name', 'unnamed')}"
Expand Down
31 changes: 20 additions & 11 deletions src/tetra_rp/core/resources/network_volume.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from pydantic import (
Field,
field_serializer,
)

from ..api.runpod import RunpodRestClient
Expand All @@ -20,8 +21,6 @@ class DataCenter(str, Enum):
"""

EU_RO_1 = "EU-RO-1"
US_WA_1 = "US-WA-1"
US_CA_1 = "US-CA-1"


class NetworkVolume(DeployableResource):
Expand All @@ -33,10 +32,20 @@ class NetworkVolume(DeployableResource):

"""

dataCenterId: Optional[DataCenter] = None
# Internal fixed value
dataCenterId: DataCenter = Field(default=DataCenter.EU_RO_1, frozen=True)

id: Optional[str] = Field(default=None)
name: Optional[str] = None
size: Optional[int] = None # Size in GB
size: Optional[int] = Field(default=10, gt=0) # Size in GB

def __str__(self) -> str:
return f"{self.__class__.__name__}:{self.id}"

@field_serializer("dataCenterId")
def serialize_data_center_id(self, value: Optional[DataCenter]) -> Optional[str]:
"""Convert DataCenter enum to string."""
return value.value if value is not None else None

@property
def is_created(self) -> bool:
Expand Down Expand Up @@ -79,17 +88,17 @@ async def deploy(self) -> "DeployableResource":
try:
# If the resource is already deployed, return it
if self.is_deployed():
log.debug(
f"Network volume {self.id} is already deployed. Mounting existing volume."
)
log.info(f"Mounted existing network volume: {self.id}")
log.debug(f"{self} exists")
return self

# Create the network volume
self = await self.create_network_volume()
async with RunpodRestClient() as client:
# Create the network volume
payload = self.model_dump(exclude_none=True)
result = await client.create_network_volume(payload)

if self.is_deployed():
return self
if volume := self.__class__(**result):
return volume

raise ValueError("Deployment failed, no volume was created.")

Expand Down
52 changes: 42 additions & 10 deletions src/tetra_rp/core/resources/serverless.py
Original file line number Diff line number Diff line change
@@ -1,27 +1,27 @@
import asyncio
import logging
from typing import Any, Dict, List, Optional
from enum import Enum
from typing import Any, Dict, List, Optional

from pydantic import (
BaseModel,
Field,
field_serializer,
field_validator,
model_validator,
BaseModel,
Field,
)

from runpod.endpoint.runner import Job

from ..api.runpod import RunpodGraphQLClient
from ..utils.backoff import get_backoff_delay

from .cloud import runpod
from .base import DeployableResource
from .template import PodTemplate, KeyValuePair
from .gpu import GpuGroup
from .cloud import runpod
from .constants import CONSOLE_URL
from .cpu import CpuInstanceType
from .environment import EnvironmentVars
from .constants import CONSOLE_URL
from .gpu import GpuGroup
from .network_volume import NetworkVolume
from .template import KeyValuePair, PodTemplate


# Environment variables are loaded from the .env file
Expand Down Expand Up @@ -62,7 +62,15 @@ class ServerlessResource(DeployableResource):
Base class for GPU serverless resource
"""

_input_only = {"id", "cudaVersions", "env", "gpus", "flashboot", "imageName"}
_input_only = {
"id",
"cudaVersions",
"env",
"gpus",
"flashboot",
"imageName",
"networkVolume",
}

# === Input-only Fields ===
cudaVersions: Optional[List[CudaVersion]] = [] # for allowedCudaVersions
Expand All @@ -71,6 +79,8 @@ class ServerlessResource(DeployableResource):
gpus: Optional[List[GpuGroup]] = [GpuGroup.ANY] # for gpuIds
imageName: Optional[str] = "" # for template.imageName

networkVolume: Optional[NetworkVolume] = None

# === Input Fields ===
executionTimeoutMs: Optional[int] = None
gpuCount: Optional[int] = 1
Expand Down Expand Up @@ -142,6 +152,10 @@ def sync_input_fields(self):
if self.flashboot:
self.name += "-fb"

if self.networkVolume and self.networkVolume.is_created:
# Volume already exists, use its ID
self.networkVolumeId = self.networkVolume.id

if self.instanceIds:
return self._sync_input_fields_cpu()
else:
Expand Down Expand Up @@ -177,6 +191,21 @@ def _sync_input_fields_cpu(self):

return self

async def _ensure_network_volume_deployed(self) -> None:
"""
Ensures network volume is deployed and ready.
Updates networkVolumeId with the deployed volume ID.
"""
Comment thread
pandyamarut marked this conversation as resolved.
if self.networkVolumeId:
return

if not self.networkVolume:
log.info(f"{self.name} requires a default network volume")
self.networkVolume = NetworkVolume(name=f"{self.name}-volume")

if deployedNetworkVolume := await self.networkVolume.deploy():
self.networkVolumeId = deployedNetworkVolume.id

def is_deployed(self) -> bool:
"""
Checks if the serverless resource is deployed and available.
Expand All @@ -202,6 +231,9 @@ async def deploy(self) -> "DeployableResource":
log.debug(f"{self} exists")
return self

# NEW: Ensure network volume is deployed first
await self._ensure_network_volume_deployed()

async with RunpodGraphQLClient() as client:
payload = self.model_dump(exclude=self._input_only, exclude_none=True)
result = await client.create_endpoint(payload)
Expand Down
Loading