diff --git a/src/tetra_rp/client.py b/src/tetra_rp/client.py index 12c5d14c..a61c4636 100644 --- a/src/tetra_rp/client.py +++ b/src/tetra_rp/client.py @@ -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 @@ -12,7 +12,6 @@ def remote( resource_config: ServerlessResource, dependencies: List[str] = None, system_dependencies: List[str] = None, - mount_volume: Optional[NetworkVolume] = None, **extra, ): """ @@ -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 diff --git a/src/tetra_rp/core/api/runpod.py b/src/tetra_rp/core/api/runpod.py index c5d94347..7524e623 100644 --- a/src/tetra_rp/core/api/runpod.py +++ b/src/tetra_rp/core/api/runpod.py @@ -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__) @@ -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')}" diff --git a/src/tetra_rp/core/resources/network_volume.py b/src/tetra_rp/core/resources/network_volume.py index c4f6e095..1b2f7285 100644 --- a/src/tetra_rp/core/resources/network_volume.py +++ b/src/tetra_rp/core/resources/network_volume.py @@ -4,6 +4,7 @@ from pydantic import ( Field, + field_serializer, ) from ..api.runpod import RunpodRestClient @@ -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): @@ -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: @@ -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.") diff --git a/src/tetra_rp/core/resources/serverless.py b/src/tetra_rp/core/resources/serverless.py index 51267fab..75c28684 100644 --- a/src/tetra_rp/core/resources/serverless.py +++ b/src/tetra_rp/core/resources/serverless.py @@ -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 @@ -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 @@ -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 @@ -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: @@ -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. + """ + 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. @@ -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)