4242from ai .backend .common .metrics .http import build_api_metric_middleware
4343from ai .backend .common .middlewares .exception import general_exception_middleware
4444from ai .backend .common .typed_validators import PydanticJWTValidator
45- from ai .backend .common .types import BinarySize , VFolderID
45+ from ai .backend .common .types import BinarySize , TusSessionId , VFolderID
4646from ai .backend .logging import BraceStyleAdapter
4747from ai .backend .storage import __version__
4848from ai .backend .storage .dto .context import StorageRootCtx
49- from ai .backend .storage .errors import InvalidAPIParameters , UploadOffsetMismatchError
49+ from ai .backend .storage .errors import (
50+ InvalidAPIParameters ,
51+ UploadChunkExceedsTotalSizeError ,
52+ UploadOffsetMismatchError ,
53+ )
5054from ai .backend .storage .services .file_stream .zip import (
5155 ZipArchiveStreamReader ,
5256)
53- from ai .backend .storage .services .upload .tus_session import (
54- TusUploadSession ,
55- TusUploadSessionArgs ,
56- )
57+ from ai .backend .storage .services .upload .tus_session import TusUploadSession
58+ from ai .backend .storage .services .upload .types import TusUploadSessionArgs
5759from ai .backend .storage .types import SENTINEL , TusChunkUploadStreamReader
5860from ai .backend .storage .utils import (
5961 CheckParamSource ,
6365
6466if TYPE_CHECKING :
6567 from ai .backend .storage .context import RootContext
66- from ai .backend .storage .volumes .abc import AbstractVolume
6768
6869log = BraceStyleAdapter (logging .getLogger (__spec__ .name ))
6970
@@ -99,7 +100,7 @@ class UploadTokenData(TypedDict):
99100 volume : str
100101 vfid : VFolderID
101102 relpath : str
102- session : str
103+ session : TusSessionId
103104 size : int
104105
105106
@@ -332,7 +333,9 @@ class Params(TypedDict):
332333 ) as params :
333334 token_data = params ["token" ]
334335 async with ctx .get_volume (token_data ["volume" ]) as volume :
335- session_dir = _resolve_tus_upload_session_dir (volume , token_data )
336+ session_dir = (
337+ volume .mangle_vfpath (token_data ["vfid" ]) / ".upload" / token_data ["session" ]
338+ )
336339 if not session_dir .exists ():
337340 raise web .HTTPNotFound (
338341 body = dump_json_str (
@@ -349,10 +352,11 @@ class Params(TypedDict):
349352 session_id = token_data ["session" ],
350353 total_size = int (token_data ["size" ]),
351354 valkey_client = ctx .valkey_tus_client ,
355+ lock_factory = ctx .tus_lock_factory ,
352356 )
353357 )
354358 state = await session .read_state ()
355- headers = _tus_response_headers (
359+ headers = _prepare_tus_session_headers (
356360 upload_offset = state .committed_offset ,
357361 upload_length = int (token_data ["size" ]),
358362 )
@@ -411,7 +415,9 @@ class Params(TypedDict):
411415 )
412416
413417 async with ctx .get_volume (token_data ["volume" ]) as volume :
414- session_dir = _resolve_tus_upload_session_dir (volume , token_data )
418+ session_dir = (
419+ volume .mangle_vfpath (token_data ["vfid" ]) / ".upload" / token_data ["session" ]
420+ )
415421 if not session_dir .exists ():
416422 raise web .HTTPNotFound (
417423 body = dump_json_str (
@@ -428,31 +434,30 @@ class Params(TypedDict):
428434 session_id = token_data ["session" ],
429435 total_size = total_size ,
430436 valkey_client = ctx .valkey_tus_client ,
437+ lock_factory = ctx .tus_lock_factory ,
431438 )
432439 )
433440 await session .ensure_initialized ()
434441
435442 upload_stream = TusChunkUploadStreamReader (
436443 request .content , request .content_type , DEFAULT_CHUNK_SIZE
437444 )
438- temp_chunk , length , sha256 = await session .write_temp_chunk (
439- client_offset , upload_stream
440- )
445+ written = await session .write_temp_chunk (client_offset , upload_stream )
441446 try :
442- if client_offset + length > total_size :
443- raise UploadOffsetMismatchError (
444- f"Chunk at offset { client_offset } with length { length } "
447+ if client_offset + written . length > total_size :
448+ raise UploadChunkExceedsTotalSizeError (
449+ f"Chunk at offset { client_offset } with length { written . length } "
445450 f"exceeds declared size { total_size } "
446451 )
447452 acceptance = await session .commit_chunk (
448453 offset = client_offset ,
449- chunk_path = temp_chunk .path ,
450- length = length ,
451- sha256 = sha256 ,
454+ chunk_path = written .path ,
455+ length = written . length ,
456+ sha256 = written . sha256 ,
452457 )
453458 except BaseException :
454- if temp_chunk .path .exists ():
455- await asyncio .to_thread (temp_chunk .path .unlink )
459+ if written .path .exists ():
460+ await asyncio .to_thread (written .path .unlink )
456461 raise
457462
458463 state = acceptance .state
@@ -464,21 +469,17 @@ class Params(TypedDict):
464469 await session .assemble (target_path )
465470 await session .cleanup ()
466471
467- headers = _tus_response_headers (
472+ headers = _prepare_tus_session_headers (
468473 upload_offset = state .committed_offset ,
469474 upload_length = total_size ,
470475 )
471476 return web .Response (status = HTTPStatus .NO_CONTENT , headers = headers )
472477
473478
474- def _resolve_tus_upload_session_dir (volume : AbstractVolume , token_data : UploadTokenData ) -> Path :
475- return volume .mangle_vfpath (token_data ["vfid" ]) / ".upload" / token_data ["session" ]
476-
477-
478479_TUS_HEADER_LIST = "Tus-Resumable, Upload-Length, Upload-Metadata, Upload-Offset, Content-Type"
479480
480481
481- def _tus_response_headers (* , upload_offset : int , upload_length : int ) -> dict [str , str ]:
482+ def _prepare_tus_session_headers (* , upload_offset : int , upload_length : int ) -> dict [str , str ]:
482483 return {
483484 "Access-Control-Allow-Origin" : "*" ,
484485 "Access-Control-Allow-Headers" : _TUS_HEADER_LIST ,
0 commit comments