This is an automated email from the ASF dual-hosted git repository. github-merge-queue[bot] pushed a commit to branch gh-readonly-queue/main/pr-6872-f236751dd601159bb14567ef4fe2509803cda809 in repository https://gitbox.apache.org/repos/asf/texera.git
commit 7fbf64e3da96a271245c3be47a2587d5513dee6f Author: Tanishq Gandhi <[email protected]> AuthorDate: Wed Aug 26 00:34:29 2026 +0000 feat(file-service): add model file upload and version endpoints (#6872) ### What changes were proposed in this PR? Adds file upload and versioning to the model API, on top of the metadata/access layer. A model version is a folder of files: uploads are staged, then committed together as one immutable LakeFS version. **Endpoints** on `ModelResource`: | Method | Path | | | --- | --- | --- | | `POST` | `/{mid}/version/create` | commit staged files as a version | | `GET` | `/{mid}/version/list` | | | `GET` | `/{mid}/version/latest` | | | `GET` | `/{mid}/version/{mvid}/rootFileNodes` | the version's file tree | | `POST` | `/{mid}/upload` | one-shot upload | | `DELETE` | `/{mid}/file` | drop a staged file before commit | | `POST` | `/multipart-upload` | init / finish / abort | | `POST` | `/multipart-upload/part` | | **Schema:** `model_upload_session` and `model_upload_session_part`, in `texera_ddl.sql` and migration `sql/updates/41.sql` (changeSet 41). Re-running the migration is a no-op. **Upload engine:** the multipart flow reuses `ResourceUploadService` from #7764 rather than a second copy of the dataset engine, so what lands here is the model descriptor plus the endpoints above. **File types:** any file is accepted. A model bundles weights with `config.json`, tokenizer/vocab files and sharded checkpoints, so an extension allowlist would reject valid models. `framework` stays metadata. ### Any related issues, documentation, discussions? Part of #6498. Umbrella #6494. Design discussion #6616. The remaining #6498 work — the endpoints the model UI needs — follows in a separate PR. ### How was this PR tested? New `ModelUploadResourceSpec` (9 tests): one-shot upload committed into a version, `.pth` accepted, a version with no staged changes rejected, staged-file delete, companion files committed alongside weights, nested folder structure preserved in the committed tree, a later version carrying over untouched files and replacing only the re-uploaded one, and the multipart init → part → finish path plus abort. Full `FileService` suite passes: **332 tests, 14 suites, 0 failures**, including `DatasetResourceSpec` (152) unchanged — the shared engine is not regressed by the model descriptor. ``` sbt "FileService/test" ``` `scalafmtCheckAll` and `scalafixAll --check` clean. ### Was this PR authored or co-authored using generative AI tooling? Generated-by: Claude Code (Claude Opus 4.8) --------- Co-authored-by: ali <[email protected]> Co-authored-by: Claude Opus 4.8 <[email protected]> --- .../texera/service/resource/DatasetResource.scala | 79 +--- .../texera/service/resource/ModelResource.scala | 410 ++++++++++++++++- .../service/resource/ResourceUploadService.scala | 77 +++- .../service/resource/ModelUploadResourceSpec.scala | 486 +++++++++++++++++++++ sql/changelog.xml | 5 + sql/texera_ddl.sql | 53 ++- sql/updates/41.sql | 82 ++++ 7 files changed, 1129 insertions(+), 63 deletions(-) diff --git a/file-service/src/main/scala/org/apache/texera/service/resource/DatasetResource.scala b/file-service/src/main/scala/org/apache/texera/service/resource/DatasetResource.scala index e9a78f3940..9f34dae154 100644 --- a/file-service/src/main/scala/org/apache/texera/service/resource/DatasetResource.scala +++ b/file-service/src/main/scala/org/apache/texera/service/resource/DatasetResource.scala @@ -504,7 +504,7 @@ class DatasetResource extends LazyLogging { LakeFSFileNode .fromLakeFSRepositoryCommittedObjects( resourceType, - Map((user.getEmail, datasetName, newVersionName) -> fileNodes) + Map((getOwner(ctx, did).getEmail, datasetName, newVersionName) -> fileNodes) ) ) } @@ -1030,37 +1030,18 @@ class DatasetResource extends LazyLogging { throw new NotFoundException(ERR_DATASET_VERSION_NOT_FOUND_MESSAGE) ) - val datasetsNode = LakeFSFileNode - .fromLakeFSRepositoryCommittedObjects( - resourceType, - Map( - ( - getOwner(ctx, did).getEmail, - dataset.getName, - latestVersion.getName - ) -> LakeFSStorageClient - .retrieveObjectsOfVersion(dataset.getRepositoryName, latestVersion.getVersionHash) - ) - ) - .head - - val ownerNode = datasetsNode.getChildren.headOption.getOrElse( - throw new IllegalStateException( - s"Dataset file tree for ${dataset.getName} is missing its owner node" - ) - ) - DashboardDatasetVersion( latestVersion, - ownerNode.children.get - .find(_.getName == dataset.getName) - .head - .children - .get - .find(_.getName == latestVersion.getName) - .head - .children - .get + ResourceUploadService + .versionRootFileNodes( + resourceType, + getOwner(ctx, did).getEmail, + dataset.getName, + latestVersion.getName, + dataset.getRepositoryName, + latestVersion.getVersionHash + ) + ._1 ) }) } @@ -1239,37 +1220,15 @@ class DatasetResource extends LazyLogging { ): DatasetVersionRootFileNodesResponse = { val dataset = getDashboardDataset(ctx, did, uid) val datasetVersion = getDatasetVersionByID(ctx, dvid) - val datasetName = dataset.dataset.getName - val repositoryName = dataset.dataset.getRepositoryName - - val datasetsNode = LakeFSFileNode - .fromLakeFSRepositoryCommittedObjects( - resourceType, - Map( - (dataset.ownerEmail, datasetName, datasetVersion.getName) -> LakeFSStorageClient - .retrieveObjectsOfVersion(repositoryName, datasetVersion.getVersionHash) - ) - ) - .head - - val ownerFileNode = datasetsNode.getChildren.headOption.getOrElse( - throw new IllegalStateException( - s"Dataset file tree for $datasetName is missing its owner node" - ) - ) - - DatasetVersionRootFileNodesResponse( - ownerFileNode.children.get - .find(_.getName == datasetName) - .head - .children - .get - .find(_.getName == datasetVersion.getName) - .head - .children - .get, - LakeFSFileNode.calculateTotalSize(List(datasetsNode)) + val (nodes, size) = ResourceUploadService.versionRootFileNodes( + resourceType, + dataset.ownerEmail, + dataset.dataset.getName, + datasetVersion.getName, + dataset.dataset.getRepositoryName, + datasetVersion.getVersionHash ) + DatasetVersionRootFileNodesResponse(nodes, size) } private def generatePresignedResponse( diff --git a/file-service/src/main/scala/org/apache/texera/service/resource/ModelResource.scala b/file-service/src/main/scala/org/apache/texera/service/resource/ModelResource.scala index e5e42563a9..ec7c1ac774 100644 --- a/file-service/src/main/scala/org/apache/texera/service/resource/ModelResource.scala +++ b/file-service/src/main/scala/org/apache/texera/service/resource/ModelResource.scala @@ -24,6 +24,7 @@ import io.dropwizard.auth.Auth import jakarta.annotation.security.{PermitAll, RolesAllowed} import jakarta.ws.rs._ import jakarta.ws.rs.core._ +import org.apache.texera.amber.core.storage.ResourceType import org.apache.texera.amber.core.storage.util.LakeFSStorageClient import org.apache.texera.auth.SessionUser import org.apache.texera.common.config.StorageConfig @@ -31,8 +32,11 @@ import org.apache.texera.dao.SqlServer import org.apache.texera.dao.SqlServer.withTransaction import org.apache.texera.dao.jooq.generated.enums.PrivilegeEnum import org.apache.texera.dao.jooq.generated.tables.Model.MODEL +import org.apache.texera.dao.jooq.generated.tables.ModelVersion.MODEL_VERSION +import org.apache.texera.dao.jooq.generated.tables.User.USER import org.apache.texera.dao.jooq.generated.tables.daos.{ModelDao, ModelUserAccessDao} -import org.apache.texera.dao.jooq.generated.tables.pojos.{Model, ModelUserAccess} +import org.apache.texera.dao.jooq.generated.tables.pojos.{Model, ModelUserAccess, ModelVersion} +import org.apache.texera.service.`type`.LakeFSFileNode import org.apache.texera.service.resource.ResourceTables.{Model => MODEL_RESOURCE} import org.apache.texera.service.resource.ModelAccessResource._ import org.apache.texera.service.resource.ModelResource.{context, _} @@ -40,11 +44,21 @@ import org.apache.texera.service.util.S3StorageClient import org.apache.texera.service.util.LakeFSExceptionHandler.withLakeFSErrorHandling import org.jooq.{DSLContext, EnumType} +import java.io.InputStream +import java.util.Optional +import scala.jdk.CollectionConverters._ +import scala.jdk.OptionConverters._ + object ModelResource { // MVP supports a single framework; stored on the model so later frameworks can be added. private val DEFAULT_FRAMEWORK = "pytorch" + // Matches model_version.name VARCHAR(128). + private val MAX_VERSION_NAME_LENGTH = 128 + + private val MULTIPART_OPERATIONS = Seq("list", "init", "finish", "abort") + private def context = SqlServer .getInstance() @@ -62,6 +76,36 @@ object ModelResource { model } + /** + * Helper function to get the model version from DB using mvid + */ + /** Scoped to `mid`: the access check runs against `mid`, so an unscoped lookup would + * resolve another model's version through this repository. + */ + private def getModelVersionByID(ctx: DSLContext, mid: Integer, mvid: Integer): ModelVersion = { + val version = ctx + .selectFrom(MODEL_VERSION) + .where(MODEL_VERSION.MVID.eq(mvid).and(MODEL_VERSION.MID.eq(mid))) + .fetchOneInto(classOf[ModelVersion]) + if (version == null) { + throw new NotFoundException("Model Version not found") + } + version + } + + /** + * Helper function to get the latest model version from the DB + */ + private def getLatestModelVersion(ctx: DSLContext, mid: Integer): Option[ModelVersion] = { + ctx + .selectFrom(MODEL_VERSION) + .where(MODEL_VERSION.MID.eq(mid)) + .orderBy(MODEL_VERSION.CREATION_TIME.desc()) + .limit(1) + .fetchOptionalInto(classOf[ModelVersion]) + .toScala + } + case class DashboardModel( model: Model, ownerEmail: String, @@ -82,12 +126,23 @@ object ModelResource { case class ModelDescriptionModification(mid: Integer, description: String) case class ModelNameModification(mid: Integer, name: String) + + case class DashboardModelVersion( + modelVersion: ModelVersion, + fileNodes: List[LakeFSFileNode] + ) + + case class ModelVersionRootFileNodesResponse( + fileNodes: List[LakeFSFileNode], + size: Long + ) } @Produces(Array(MediaType.APPLICATION_JSON)) @Path("/model") class ModelResource extends LazyLogging { private val ERR_USER_HAS_NO_ACCESS_TO_MODEL_MESSAGE = "User has no access to this model" + private val ERR_MODEL_VERSION_NOT_FOUND_MESSAGE = "The version of the model not found" /** * Helper function to get the model from DB with additional information including @@ -408,4 +463,357 @@ class ModelResource extends LazyLogging { ): DashboardModel = { withTransaction(context)(ctx => getDashboardModel(ctx, mid, None)) } + + // =========================================================================== + // Versioning + // =========================================================================== + + @POST + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/version/create") + @Consumes(Array(MediaType.TEXT_PLAIN)) + def createModelVersion( + versionName: String, + @PathParam("mid") mid: Integer, + @Auth user: SessionUser + ): DashboardModelVersion = { + val uid = user.getUid + withTransaction(context) { ctx => + if (!userHasWriteAccess(ctx, mid, uid)) { + throw new ForbiddenException(ERR_USER_HAS_NO_ACCESS_TO_MODEL_MESSAGE) + } + + val model = getModelByID(ctx, mid) + val modelName = model.getName + val repositoryName = model.getRepositoryName + + // Check if there are any changes in LakeFS before creating a new version + val diffs = withLakeFSErrorHandling { + LakeFSStorageClient.retrieveUncommittedObjects(repoName = repositoryName) + } + + if (diffs.isEmpty) { + throw new WebApplicationException( + "No changes detected in model. Version creation aborted.", + Response.Status.BAD_REQUEST + ) + } + + // Generate a new version name + val versionCount = ctx + .selectCount() + .from(MODEL_VERSION) + .where(MODEL_VERSION.MID.eq(mid)) + .fetchOne(0, classOf[Int]) + + val sanitizedVersionName = Option(versionName).filter(_.nonEmpty).getOrElse("") + val newVersionName = if (sanitizedVersionName.isEmpty) { + s"v${versionCount + 1}" + } else { + s"v${versionCount + 1} - $sanitizedVersionName" + } + + // Before the commit: the commit is outside this transaction, so a name the insert + // rejects would leave a commit no version points at and strand the staged file. + if (newVersionName.length > MAX_VERSION_NAME_LENGTH) { + throw new BadRequestException( + s"Version name is too long: ${newVersionName.length} characters, " + + s"maximum is $MAX_VERSION_NAME_LENGTH." + ) + } + + // Create a commit in LakeFS + val commit = withLakeFSErrorHandling { + LakeFSStorageClient.createCommit( + repoName = repositoryName, + branch = "main", + commitMessage = s"Created model version: $newVersionName" + ) + } + + if (commit == null || commit.getId == null) { + throw new WebApplicationException( + "Failed to create commit in LakeFS. Version creation aborted.", + Response.Status.INTERNAL_SERVER_ERROR + ) + } + + // Create a new model version entry in the database + val modelVersion = new ModelVersion() + modelVersion.setMid(mid) + modelVersion.setCreatorUid(uid) + modelVersion.setName(newVersionName) + modelVersion.setVersionHash(commit.getId) // Store LakeFS version hash + + val insertedVersion = ctx + .insertInto(MODEL_VERSION) + .set(ctx.newRecord(MODEL_VERSION, modelVersion)) + .returning() + .fetchOne() + .into(classOf[ModelVersion]) + + // Retrieve committed file structure + val fileNodes = withLakeFSErrorHandling { + LakeFSStorageClient.retrieveObjectsOfVersion(repositoryName, commit.getId) + } + + DashboardModelVersion( + insertedVersion, + LakeFSFileNode + .fromLakeFSRepositoryCommittedObjects( + ResourceType.Model, + Map((getOwner(ctx, mid).getEmail, modelName, newVersionName) -> fileNodes) + ) + ) + } + } + + @GET + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/version/list") + def getModelVersionList( + @PathParam("mid") mid: Integer, + @Auth user: SessionUser + ): List[ModelVersion] = { + val uid = user.getUid + withTransaction(context)(ctx => { + val model = getModelByID(ctx, mid) + if (!userHasReadAccess(ctx, model.getMid, uid)) { + throw new ForbiddenException(ERR_USER_HAS_NO_ACCESS_TO_MODEL_MESSAGE) + } + fetchModelVersions(ctx, model.getMid) + }) + } + + @GET + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/version/latest") + def retrieveLatestModelVersion( + @PathParam("mid") mid: Integer, + @Auth user: SessionUser + ): DashboardModelVersion = { + val uid = user.getUid + withTransaction(context)(ctx => { + if (!userHasReadAccess(ctx, mid, uid)) { + throw new ForbiddenException(ERR_USER_HAS_NO_ACCESS_TO_MODEL_MESSAGE) + } + val latestVersion = getLatestModelVersion(ctx, mid).getOrElse( + throw new NotFoundException(ERR_MODEL_VERSION_NOT_FOUND_MESSAGE) + ) + DashboardModelVersion(latestVersion, versionRootFileNodes(ctx, mid, latestVersion)) + }) + } + + @GET + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/version/{mvid}/rootFileNodes") + def retrieveModelVersionRootFileNodes( + @PathParam("mid") mid: Integer, + @PathParam("mvid") mvid: Integer, + @Auth user: SessionUser + ): ModelVersionRootFileNodesResponse = { + val uid = user.getUid + withTransaction(context)(ctx => fetchModelVersionRootFileNodes(ctx, mid, mvid, Some(uid))) + } + + // =========================================================================== + // File upload (one-shot + session-based multipart) + // =========================================================================== + + @POST + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/upload") + @Consumes(Array(MediaType.APPLICATION_OCTET_STREAM)) + def uploadOneFileToModel( + @PathParam("mid") mid: Integer, + @QueryParam("filePath") encodedFilePath: String, + @QueryParam("message") message: String, + fileStream: InputStream, + @Context headers: HttpHeaders, + @Auth user: SessionUser + ): Response = { + ResourceUploadService.uploadOneFile( + ResourceStorage.Model, + mid, + encodedFilePath, + fileStream, + headers, + user.getUid + ) + } + + @DELETE + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/{mid}/file") + @Consumes(Array(MediaType.APPLICATION_JSON)) + def deleteModelFile( + @PathParam("mid") mid: Integer, + @QueryParam("filePath") encodedFilePath: String, + @Auth user: SessionUser + ): Response = { + ResourceUploadService.deleteStagedFile( + ResourceStorage.Model, + mid, + encodedFilePath, + user.getUid + ) + } + + @POST + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Path("/multipart-upload") + @Consumes(Array(MediaType.APPLICATION_JSON)) + def multipartUpload( + @QueryParam("type") operationType: String, + @QueryParam("ownerEmail") ownerEmail: String, + @QueryParam("modelName") modelName: String, + @QueryParam("filePath") filePath: String, + @QueryParam("fileSizeBytes") fileSizeBytes: Optional[java.lang.Long], + @QueryParam("partSizeBytes") partSizeBytes: Optional[java.lang.Long], + @QueryParam("restart") restart: Optional[java.lang.Boolean], + @Auth user: SessionUser + ): Response = { + val uid = user.getUid + + // Optional query param: null when omitted, so validate before dereferencing and before + // the getModelBy round-trip. + val operation = Option(operationType).map(_.trim.toLowerCase).getOrElse("") + if (!MULTIPART_OPERATIONS.contains(operation)) { + throw new BadRequestException( + s"Invalid type parameter. Use ${MULTIPART_OPERATIONS.map(o => s"'$o'").mkString(", ")}." + ) + } + + val model: Model = getModelBy(ownerEmail, modelName) + + operation match { + case "list" => listMultipartUploads(model.getMid, uid) + case "init" => + initMultipartUpload(model.getMid, filePath, fileSizeBytes, partSizeBytes, restart, uid) + case "finish" => finishMultipartUpload(model.getMid, filePath, uid) + case _ => abortMultipartUpload(model.getMid, filePath, uid) + } + } + + @POST + @RolesAllowed(Array("REGULAR", "ADMIN")) + @Consumes(Array(MediaType.APPLICATION_OCTET_STREAM)) + @Path("/multipart-upload/part") + def uploadPart( + @QueryParam("ownerEmail") modelOwnerEmail: String, + @QueryParam("modelName") modelName: String, + @QueryParam("filePath") encodedFilePath: String, + @QueryParam("partNumber") partNumber: Int, + partStream: InputStream, + @Context headers: HttpHeaders, + @Auth user: SessionUser + ): Response = { + val model = getModelBy(modelOwnerEmail, modelName) + ResourceUploadService.uploadPart( + ResourceStorage.Model, + model.getMid, + user.getUid, + encodedFilePath, + partNumber, + partStream, + headers + ) + } + + // =========================================================================== + // Private helpers + // =========================================================================== + + private def fetchModelVersions(ctx: DSLContext, mid: Integer): List[ModelVersion] = { + ctx + .selectFrom(MODEL_VERSION) + .where(MODEL_VERSION.MID.eq(mid)) + .orderBy(MODEL_VERSION.CREATION_TIME.desc()) + .fetchInto(classOf[ModelVersion]) + .asScala + .toList + } + + /** + * Builds the file-tree children of a single model version, drilling into the + * owner/model/version nesting produced by LakeFSFileNode. + */ + private def versionRootFileNodes( + ctx: DSLContext, + mid: Integer, + modelVersion: ModelVersion + ): List[LakeFSFileNode] = { + val model = getModelByID(ctx, mid) + ResourceUploadService + .versionRootFileNodes( + ResourceType.Model, + getOwner(ctx, mid).getEmail, + model.getName, + modelVersion.getName, + model.getRepositoryName, + modelVersion.getVersionHash + ) + ._1 + } + + private def fetchModelVersionRootFileNodes( + ctx: DSLContext, + mid: Integer, + mvid: Integer, + uid: Option[Integer] + ): ModelVersionRootFileNodesResponse = { + val model = getDashboardModel(ctx, mid, uid) + val modelVersion = getModelVersionByID(ctx, mid, mvid) + val (nodes, size) = ResourceUploadService.versionRootFileNodes( + ResourceType.Model, + model.ownerEmail, + model.model.getName, + modelVersion.getName, + model.model.getRepositoryName, + modelVersion.getVersionHash + ) + ModelVersionRootFileNodesResponse(nodes, size) + } + + private def getModelBy(ownerEmail: String, modelName: String): Model = { + val model = context + .select(MODEL.fields: _*) + .from(MODEL) + .leftJoin(USER) + .on(USER.UID.eq(MODEL.OWNER_UID)) + .where(USER.EMAIL.eq(ownerEmail)) + .and(MODEL.NAME.eq(modelName)) + .fetchOneInto(classOf[Model]) + if (model == null) { + throw new BadRequestException("Model not found") + } + model + } + + private def listMultipartUploads(mid: Integer, requesterUid: Int): Response = + ResourceUploadService.listUploads(ResourceStorage.Model, mid, requesterUid) + + private def initMultipartUpload( + mid: Integer, + encodedFilePath: String, + fileSizeBytes: Optional[java.lang.Long], + partSizeBytes: Optional[java.lang.Long], + restart: Optional[java.lang.Boolean], + uid: Integer + ): Response = + ResourceUploadService.initUpload( + ResourceStorage.Model, + mid, + encodedFilePath, + fileSizeBytes, + partSizeBytes, + restart, + uid + ) + + private def finishMultipartUpload(mid: Integer, encodedFilePath: String, uid: Int): Response = + ResourceUploadService.finishUpload(ResourceStorage.Model, mid, encodedFilePath, uid) + + private def abortMultipartUpload(mid: Integer, encodedFilePath: String, uid: Int): Response = + ResourceUploadService.abortUpload(ResourceStorage.Model, mid, encodedFilePath, uid) } diff --git a/file-service/src/main/scala/org/apache/texera/service/resource/ResourceUploadService.scala b/file-service/src/main/scala/org/apache/texera/service/resource/ResourceUploadService.scala index 25c4b2a0bb..693e25c510 100644 --- a/file-service/src/main/scala/org/apache/texera/service/resource/ResourceUploadService.scala +++ b/file-service/src/main/scala/org/apache/texera/service/resource/ResourceUploadService.scala @@ -27,14 +27,22 @@ import org.apache.texera.common.config.StorageConfig import org.apache.texera.dao.{SiteSettings, SqlServer} import org.apache.texera.dao.SqlServer.withTransaction import org.apache.texera.dao.jooq.generated.tables.Dataset.DATASET +import org.apache.texera.dao.jooq.generated.tables.Model.MODEL +import org.apache.texera.dao.jooq.generated.tables.ModelUploadSession.MODEL_UPLOAD_SESSION +import org.apache.texera.dao.jooq.generated.tables.ModelUploadSessionPart.MODEL_UPLOAD_SESSION_PART import org.apache.texera.dao.jooq.generated.tables.DatasetUploadSession.DATASET_UPLOAD_SESSION import org.apache.texera.dao.jooq.generated.tables.DatasetUploadSessionPart.DATASET_UPLOAD_SESSION_PART import org.apache.texera.dao.jooq.generated.tables.records.{ DatasetRecord, DatasetUploadSessionPartRecord, DatasetUploadSessionRecord, - DatasetUserAccessRecord + DatasetUserAccessRecord, + ModelRecord, + ModelUploadSessionPartRecord, + ModelUploadSessionRecord, + ModelUserAccessRecord } +import org.apache.texera.service.`type`.LakeFSFileNode import org.apache.texera.service.util.LakeFSExceptionHandler.withLakeFSErrorHandling import org.apache.texera.service.util.S3StorageClient import org.apache.texera.service.util.S3StorageClient.{ @@ -119,6 +127,30 @@ object ResourceStorage { partNumber = DATASET_UPLOAD_SESSION_PART.PART_NUMBER, partEtag = DATASET_UPLOAD_SESSION_PART.ETAG ) + + val Model: ResourceStorage[ + ModelRecord, + ModelUserAccessRecord, + ModelUploadSessionRecord, + ModelUploadSessionPartRecord + ] = + ResourceStorage( + resource = ResourceTables.Model, + resourceType = ResourceType.Model, + repositoryNameField = MODEL.REPOSITORY_NAME, + sessionResourceId = MODEL_UPLOAD_SESSION.MID, + sessionUid = MODEL_UPLOAD_SESSION.UID, + sessionFilePath = MODEL_UPLOAD_SESSION.FILE_PATH, + sessionUploadId = MODEL_UPLOAD_SESSION.UPLOAD_ID, + sessionPhysicalAddress = MODEL_UPLOAD_SESSION.PHYSICAL_ADDRESS, + sessionNumParts = MODEL_UPLOAD_SESSION.NUM_PARTS_REQUESTED, + sessionFileSize = MODEL_UPLOAD_SESSION.FILE_SIZE_BYTES, + sessionPartSize = MODEL_UPLOAD_SESSION.PART_SIZE_BYTES, + sessionCreatedAt = MODEL_UPLOAD_SESSION.CREATED_AT, + partUploadId = MODEL_UPLOAD_SESSION_PART.UPLOAD_ID, + partNumber = MODEL_UPLOAD_SESSION_PART.PART_NUMBER, + partEtag = MODEL_UPLOAD_SESSION_PART.ETAG + ) } /** @@ -139,6 +171,49 @@ object ResourceUploadService { private def singleFileUploadMaxBytes(defaultMiB: Long = 20L): Long = SiteSettings.getLong("single_file_upload_max_size_mib", defaultMiB) * 1024L * 1024L + /** + * Builds the file nodes of one committed version, plus the version's total size. + * + * The tree is rooted at the resource-type prefix, so the paths it yields resolve against + * the right table when they are handed back to `FileResolver`. + */ + def versionRootFileNodes( + resourceType: ResourceType.Value, + ownerEmail: String, + resourceName: String, + versionName: String, + repositoryName: String, + versionHash: String + ): (List[LakeFSFileNode], Long) = { + val rootNode = LakeFSFileNode + .fromLakeFSRepositoryCommittedObjects( + resourceType, + Map( + (ownerEmail, resourceName, versionName) -> LakeFSStorageClient + .retrieveObjectsOfVersion(repositoryName, versionHash) + ) + ) + .head + + val ownerFileNode = rootNode.getChildren.headOption.getOrElse( + throw new IllegalStateException( + s"File tree for $resourceName is missing its owner node" + ) + ) + + val nodes = ownerFileNode.children.get + .find(_.getName == resourceName) + .head + .children + .get + .find(_.getName == versionName) + .head + .children + .get + + (nodes, LakeFSFileNode.calculateTotalSize(List(rootNode))) + } + private def noAccessMessage[R <: Record, A <: Record, S <: Record, P <: Record]( s: ResourceStorage[R, A, S, P] ): String = s"User has no access to this ${s.resource.label}" diff --git a/file-service/src/test/scala/org/apache/texera/service/resource/ModelUploadResourceSpec.scala b/file-service/src/test/scala/org/apache/texera/service/resource/ModelUploadResourceSpec.scala new file mode 100644 index 0000000000..82444491a2 --- /dev/null +++ b/file-service/src/test/scala/org/apache/texera/service/resource/ModelUploadResourceSpec.scala @@ -0,0 +1,486 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.texera.service.resource + +import jakarta.ws.rs._ +import jakarta.ws.rs.core._ +import org.apache.texera.auth.SessionUser +import org.apache.texera.dao.MockTexeraDB +import org.apache.texera.dao.jooq.generated.enums.{PrivilegeEnum, UserRoleEnum} +import org.apache.texera.dao.jooq.generated.tables.daos.{ModelUserAccessDao, UserDao} +import org.apache.texera.dao.jooq.generated.tables.pojos.{ModelUserAccess, User} +import org.apache.texera.service.MockLakeFS +import org.apache.texera.service.`type`.LakeFSFileNode +import org.scalatest.flatspec.AnyFlatSpec +import org.scalatest.matchers.should.Matchers +import org.scalatest.{BeforeAndAfterAll, BeforeAndAfterEach} + +import java.io.ByteArrayInputStream +import java.net.URLEncoder +import java.nio.charset.StandardCharsets +import java.util.{Collections, Date, Locale, Optional} +import scala.util.Random + +class ModelUploadResourceSpec + extends AnyFlatSpec + with Matchers + with MockTexeraDB + with MockLakeFS + with BeforeAndAfterAll + with BeforeAndAfterEach { + + private val ownerUser: User = { + val user = new User + user.setName("model_upload_user") + user.setEmail("[email protected]") + user.setRole(UserRoleEnum.ADMIN) + user + } + + /** A second account that is granted WRITE on a model but never owns one. */ + private val collaboratorUser: User = { + val user = new User + user.setName("model_upload_collaborator") + user.setEmail("[email protected]") + user.setRole(UserRoleEnum.REGULAR) + user + } + + lazy val modelResource = new ModelResource() + lazy val sessionUser = new SessionUser(ownerUser) + lazy val collaboratorSession = new SessionUser(collaboratorUser) + + override protected def beforeAll(): Unit = { + super.beforeAll() + initializeDBAndReplaceDSLContext() + val userDao = new UserDao(getDSLContext.configuration()) + userDao.insert(ownerUser) + userDao.insert(collaboratorUser) + } + + override protected def afterAll(): Unit = { + try shutdownDB() + finally super.afterAll() + } + + // ---------- helpers ---------- + private def urlEnc(raw: String): String = + URLEncoder.encode(raw, StandardCharsets.UTF_8.name()) + + private def uniqueName(prefix: String): String = + s"$prefix-${System.nanoTime()}-${Random.alphanumeric.take(6).mkString.toLowerCase}" + + /** Minimal HttpHeaders exposing only Content-Length, which the upload paths read. */ + private def mkHeaders(contentLength: Long): HttpHeaders = + new HttpHeaders { + private val headers = new MultivaluedHashMap[String, String]() + headers.putSingle(HttpHeaders.CONTENT_LENGTH, contentLength.toString) + override def getHeaderString(name: String): String = headers.getFirst(name) + override def getRequestHeaders: MultivaluedMap[String, String] = headers + override def getRequestHeader(name: String): java.util.List[String] = + Option(headers.get(name)).getOrElse(Collections.emptyList[String]()) + override def getAcceptableMediaTypes: java.util.List[MediaType] = Collections.emptyList() + override def getAcceptableLanguages: java.util.List[Locale] = Collections.emptyList() + override def getMediaType: MediaType = null + override def getLanguage: Locale = null + override def getCookies: java.util.Map[String, Cookie] = Collections.emptyMap() + override def getDate: Date = null + override def getLength: Int = contentLength.toInt + } + + /** Creates a fresh model (provisions its LakeFS repo) and returns it. */ + private def newModel(): ModelResource.DashboardModel = + modelResource.createModel( + ModelResource.CreateModelRequest( + modelName = uniqueName("upload-model"), + modelDescription = "for upload tests", + isModelPublic = false, + isModelDownloadable = true, + framework = "pytorch", + format = null + ), + sessionUser + ) + + private def uploadOneShot(mid: Integer, path: String, bytes: Array[Byte]): Response = + modelResource.uploadOneFileToModel( + mid, + urlEnc(path), + "upload", + new ByteArrayInputStream(bytes), + mkHeaders(bytes.length.toLong), + sessionUser + ) + + // =========================================================================== + // One-shot upload + version lifecycle + // =========================================================================== + "uploadOneFileToModel + createModelVersion" should "commit an uploaded .pt file into a version" in { + val model = newModel() + val mid = model.model.getMid + + uploadOneShot(mid, "model.pt", Array.fill[Byte](2048)(0x5a)).getStatus shouldEqual 200 + + val version = modelResource.createModelVersion("initial", mid, sessionUser) + version.modelVersion.getName should startWith("v1") + + modelResource.getModelVersionList(mid, sessionUser) should have size 1 + + val latest = modelResource.retrieveLatestModelVersion(mid, sessionUser) + latest.fileNodes.map(_.getName) should contain("model.pt") + + val roots = + modelResource.retrieveModelVersionRootFileNodes( + mid, + version.modelVersion.getMvid, + sessionUser + ) + roots.fileNodes.map(_.getName) should contain("model.pt") + roots.size should be > 0L + + // The serialized path must carry the "model" resource-type prefix: FileResolver + // keys on that first segment to pick the backing table, so a "/dataset/..." path + // here would resolve a model against the dataset table. + roots.fileNodes + .find(_.getName == "model.pt") + .get + .getFilePath shouldBe s"/model/${ownerUser.getEmail}/${model.model.getName}/${version.modelVersion.getName}/model.pt" + } + + it should "accept a .pth extension as well" in { + val model = newModel() + uploadOneShot( + model.model.getMid, + "weights.pth", + Array.fill[Byte](1024)(0x1) + ).getStatus shouldEqual 200 + } + + "createModelVersion" should "reject a version when there are no staged changes" in { + val model = newModel() + val ex = intercept[WebApplicationException] { + modelResource.createModelVersion("empty", model.model.getMid, sessionUser) + } + ex.getResponse.getStatus shouldEqual 400 + } + + "deleteModelFile" should "remove a staged file" in { + val model = newModel() + val mid = model.model.getMid + uploadOneShot(mid, "scratch.pt", Array.fill[Byte](512)(0x2)).getStatus shouldEqual 200 + modelResource.deleteModelFile(mid, urlEnc("scratch.pt"), sessionUser).getStatus shouldEqual 200 + } + + // =========================================================================== + // No per-file type restriction: a model is a folder of files + // =========================================================================== + "uploadOneFileToModel" should "accept companion files alongside weights and commit them together" in { + val model = newModel() + val mid = model.model.getMid + + // a typical model folder: weights plus config/tokenizer companions + uploadOneShot(mid, "model.pt", Array.fill[Byte](256)(0x5)).getStatus shouldEqual 200 + uploadOneShot( + mid, + "config.json", + "{\"hidden\":8}".getBytes(StandardCharsets.UTF_8) + ).getStatus shouldEqual 200 + uploadOneShot(mid, "tokenizer.txt", Array.fill[Byte](32)(0x4)).getStatus shouldEqual 200 + + val version = modelResource.createModelVersion("folder", mid, sessionUser) + version.fileNodes.nonEmpty shouldBe true + + val names = modelResource.retrieveLatestModelVersion(mid, sessionUser).fileNodes.map(_.getName) + names should contain allOf ("model.pt", "config.json", "tokenizer.txt") + } + + it should "preserve a nested folder structure in the committed version tree" in { + val model = newModel() + val mid = model.model.getMid + + // a HuggingFace-style layout: files inside subdirectories + uploadOneShot(mid, "pytorch_model.bin", Array.fill[Byte](128)(0x6)).getStatus shouldEqual 200 + uploadOneShot( + mid, + "tokenizer/vocab.txt", + Array.fill[Byte](64)(0x7) + ).getStatus shouldEqual 200 + uploadOneShot( + mid, + "shards/part-00001/data.bin", + Array.fill[Byte](64)(0x8) + ).getStatus shouldEqual 200 + + modelResource.createModelVersion("nested", mid, sessionUser) + + val roots = modelResource.retrieveLatestModelVersion(mid, sessionUser).fileNodes + roots.map(_.getName) should contain allOf ("pytorch_model.bin", "tokenizer", "shards") + + // directories are preserved as directory nodes holding their children + val tokenizerDir = roots.find(_.getName == "tokenizer").get + tokenizerDir.getNodeType shouldEqual "directory" + tokenizerDir.getChildren.map(_.getName) should contain("vocab.txt") + + // nesting is recursive, not flattened to one level + val shardsDir = roots.find(_.getName == "shards").get + val partDir = shardsDir.getChildren.find(_.getName == "part-00001").get + partDir.getNodeType shouldEqual "directory" + partDir.getChildren.map(_.getName) should contain("data.bin") + } + + // =========================================================================== + // Version semantics: each version is a full snapshot, not a delta + // =========================================================================== + "a later version" should "carry over untouched files and only replace the re-uploaded one" in { + val model = newModel() + val mid = model.model.getMid + + // v1: four files, with b at a known size + uploadOneShot(mid, "a.pt", Array.fill[Byte](100)(0x1)).getStatus shouldEqual 200 + uploadOneShot(mid, "b.pt", Array.fill[Byte](200)(0x2)).getStatus shouldEqual 200 + uploadOneShot(mid, "c.pt", Array.fill[Byte](300)(0x3)).getStatus shouldEqual 200 + uploadOneShot(mid, "d.pt", Array.fill[Byte](400)(0x4)).getStatus shouldEqual 200 + val v1 = modelResource.createModelVersion("first", mid, sessionUser) + + // v2: re-upload ONLY b, with a different size so the two revisions are distinguishable + uploadOneShot(mid, "b.pt", Array.fill[Byte](999)(0x9)).getStatus shouldEqual 200 + val v2 = modelResource.createModelVersion("second", mid, sessionUser) + + def nodesOf(mvid: Integer) = + modelResource.retrieveModelVersionRootFileNodes(mid, mvid, sessionUser).fileNodes + def sizeOf(mvid: Integer, name: String) = + nodesOf(mvid).find(_.getName == name).flatMap(_.getSize) + + // v2 still contains all four files: a, c, d carried over untouched, b replaced + nodesOf(v2.modelVersion.getMvid) + .map(_.getName) should contain allOf ("a.pt", "b.pt", "c.pt", "d.pt") + sizeOf(v2.modelVersion.getMvid, "a.pt") shouldEqual Some(100L) + sizeOf(v2.modelVersion.getMvid, "c.pt") shouldEqual Some(300L) + sizeOf(v2.modelVersion.getMvid, "d.pt") shouldEqual Some(400L) + sizeOf(v2.modelVersion.getMvid, "b.pt") shouldEqual Some(999L) + + // v1 is immutable: it still sees the ORIGINAL b + sizeOf(v1.modelVersion.getMvid, "b.pt") shouldEqual Some(200L) + + // both versions are listed, newest first + modelResource.getModelVersionList(mid, sessionUser).map(_.getName) should have size 2 + } + + // =========================================================================== + // Path ownership + // =========================================================================== + "a version created by a WRITE collaborator" should "carry the owner's email in its file paths" in { + val model = newModel() + val mid = model.model.getMid + + // The collaborator can write, but the model still belongs to ownerUser. + new ModelUserAccessDao(getDSLContext.configuration()) + .insert(new ModelUserAccess(mid, collaboratorUser.getUid, PrivilegeEnum.WRITE)) + + modelResource + .uploadOneFileToModel( + mid, + urlEnc("weights.bin"), + "upload", + new ByteArrayInputStream(Array.fill[Byte](64)(0x3)), + mkHeaders(64L), + collaboratorSession + ) + .getStatus shouldEqual 200 + + val created = modelResource.createModelVersion("from-collab", mid, collaboratorSession) + val versionName = created.modelVersion.getName + + // FileResolver resolves /model/<ownerEmail>/... via MODEL.OWNER_UID, so a path naming + // the collaborator resolves to nothing. + val expected = + s"/model/${ownerUser.getEmail}/${model.model.getName}/$versionName/weights.bin" + + def pathsOf(nodes: List[LakeFSFileNode]): List[String] = + nodes.flatMap(n => n.getFilePath :: pathsOf(n.getChildren)) + + pathsOf(created.fileNodes) should contain(expected) + pathsOf(created.fileNodes).foreach(_ should not include collaboratorUser.getEmail) + + // The response must agree with a subsequent read. + pathsOf( + modelResource.retrieveLatestModelVersion(mid, collaboratorSession).fileNodes + ) should contain(expected) + } + + // =========================================================================== + // Input validation and version scoping + // =========================================================================== + "createModelVersion" should "reject an over-long name without committing to LakeFS" in { + val model = newModel() + val mid = model.model.getMid + + uploadOneShot(mid, "weights.pt", Array.fill[Byte](32)(0x1)).getStatus shouldEqual 200 + + // The insert happens after the LakeFS commit, so an unchecked name would strand the + // staged file behind a commit no version points at. + val tooLong = "x" * 200 + val thrown = intercept[BadRequestException] { + modelResource.createModelVersion(tooLong, mid, sessionUser) + } + thrown.getMessage should include("too long") + + // The staged change survived, so a retry works -- previously it hit "No changes detected". + val recovered = modelResource.createModelVersion("sane-name", mid, sessionUser) + recovered.modelVersion.getName should endWith("sane-name") + modelResource + .retrieveLatestModelVersion(mid, sessionUser) + .fileNodes + .map(_.getName) should contain("weights.pt") + } + + "retrieveModelVersionRootFileNodes" should "404 for a version belonging to another model" in { + val modelA = newModel() + val modelB = newModel() + + uploadOneShot(modelA.model.getMid, "a.pt", Array.fill[Byte](16)(0x1)).getStatus shouldEqual 200 + uploadOneShot(modelB.model.getMid, "b.pt", Array.fill[Byte](16)(0x2)).getStatus shouldEqual 200 + modelResource.createModelVersion("va", modelA.model.getMid, sessionUser) + val versionOfB = modelResource.createModelVersion("vb", modelB.model.getMid, sessionUser) + + // An unscoped lookup would resolve B's version through A's repository and 500 in LakeFS. + intercept[NotFoundException] { + modelResource.retrieveModelVersionRootFileNodes( + modelA.model.getMid, + versionOfB.modelVersion.getMvid, + sessionUser + ) + } + } + + "multipartUpload" should "400 when the operation type is missing or unknown" in { + val model = newModel() + val ownerEmail = ownerUser.getEmail + val modelName = model.model.getName + + // Absent means null, which used to NPE into a 500. + for (op <- Seq(null, "", "bogus")) { + intercept[BadRequestException] { + modelResource.multipartUpload( + op, + ownerEmail, + modelName, + urlEnc("f.pt"), + Optional.empty(), + Optional.empty(), + Optional.empty(), + sessionUser + ) + } + } + } + + // =========================================================================== + // Session-based multipart upload (single part) + // =========================================================================== + "the multipart flow" should "init, upload a part, finish, and be committable as a version" in { + val model = newModel() + val mid = model.model.getMid + val ownerEmail = ownerUser.getEmail + val modelName = model.model.getName + val filePath = "multipart-model.pt" + val payload = Array.fill[Byte](16)(0x7) + val partSize = 8L * 1024L * 1024L + + // init -> one part expected + val initResp = modelResource.multipartUpload( + "init", + ownerEmail, + modelName, + urlEnc(filePath), + Optional.of(java.lang.Long.valueOf(payload.length.toLong)), + Optional.of(java.lang.Long.valueOf(partSize)), + Optional.empty(), + sessionUser + ) + initResp.getStatus shouldEqual 200 + + // upload the single part + val partResp = modelResource.uploadPart( + ownerEmail, + modelName, + urlEnc(filePath), + 1, + new ByteArrayInputStream(payload), + mkHeaders(payload.length.toLong), + sessionUser + ) + partResp.getStatus shouldEqual 200 + + // finish + val finishResp = modelResource.multipartUpload( + "finish", + ownerEmail, + modelName, + urlEnc(filePath), + Optional.empty(), + Optional.empty(), + Optional.empty(), + sessionUser + ) + finishResp.getStatus shouldEqual 200 + + // the finished file is now staged and can be committed as a version + val version = modelResource.createModelVersion("from-multipart", mid, sessionUser) + version.fileNodes.nonEmpty shouldBe true + modelResource + .retrieveLatestModelVersion(mid, sessionUser) + .fileNodes + .map(_.getName) should contain(filePath) + } + + it should "abort an initiated upload" in { + val model = newModel() + val ownerEmail = ownerUser.getEmail + val modelName = model.model.getName + val filePath = "abort-model.pt" + + modelResource + .multipartUpload( + "init", + ownerEmail, + modelName, + urlEnc(filePath), + Optional.of(java.lang.Long.valueOf(16L)), + Optional.of(java.lang.Long.valueOf(8L * 1024L * 1024L)), + Optional.empty(), + sessionUser + ) + .getStatus shouldEqual 200 + + modelResource + .multipartUpload( + "abort", + ownerEmail, + modelName, + urlEnc(filePath), + Optional.empty(), + Optional.empty(), + Optional.empty(), + sessionUser + ) + .getStatus shouldEqual 200 + } +} diff --git a/sql/changelog.xml b/sql/changelog.xml index 3debbad196..e86f869665 100644 --- a/sql/changelog.xml +++ b/sql/changelog.xml @@ -114,6 +114,11 @@ <sqlFile path="sql/updates/40.sql"/> </changeSet> + <!-- Add multipart upload session tables for models --> + <changeSet id="41" author="tanishqgandhi1908"> + <sqlFile path="sql/updates/41.sql"/> + </changeSet> + <!-- example changeSet <changeSet id="1" author="author"> <sqlFile path="sql/updates/1.sql"/> diff --git a/sql/texera_ddl.sql b/sql/texera_ddl.sql index c016817a39..acd4eef817 100644 --- a/sql/texera_ddl.sql +++ b/sql/texera_ddl.sql @@ -65,6 +65,8 @@ DROP TABLE IF EXISTS dataset_upload_session_part CASCADE; DROP TABLE IF EXISTS dataset CASCADE; DROP TABLE IF EXISTS dataset_user_access CASCADE; DROP TABLE IF EXISTS dataset_version CASCADE; +DROP TABLE IF EXISTS model_upload_session CASCADE; +DROP TABLE IF EXISTS model_upload_session_part CASCADE; DROP TABLE IF EXISTS model_user_access CASCADE; DROP TABLE IF EXISTS model_version CASCADE; DROP TABLE IF EXISTS model CASCADE; @@ -447,7 +449,9 @@ CREATE TABLE IF NOT EXISTS model_version name VARCHAR(128) NOT NULL, version_hash VARCHAR(64) NOT NULL, creation_time TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, - FOREIGN KEY (mid) REFERENCES model(mid) ON DELETE CASCADE + FOREIGN KEY (mid) REFERENCES model(mid) ON DELETE CASCADE, + -- FileResolver resolves a version by (mid, name) with fetchOneInto. + CONSTRAINT uq_model_version_mid_name UNIQUE (mid, name) ); -- model_user_access @@ -461,6 +465,53 @@ CREATE TABLE IF NOT EXISTS model_user_access FOREIGN KEY (uid) REFERENCES "user"(uid) ON DELETE CASCADE ); +-- model_upload_session +CREATE TABLE IF NOT EXISTS model_upload_session +( + mid INT NOT NULL, + uid INT NOT NULL, + file_path TEXT NOT NULL, + upload_id VARCHAR(256) NOT NULL UNIQUE, + physical_address TEXT, + num_parts_requested INT NOT NULL, + file_size_bytes BIGINT NOT NULL, + part_size_bytes BIGINT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + + PRIMARY KEY (uid, mid, file_path), + + FOREIGN KEY (mid) REFERENCES model(mid) ON DELETE CASCADE, + FOREIGN KEY (uid) REFERENCES "user"(uid) ON DELETE CASCADE, + + CONSTRAINT chk_model_upload_session_num_parts_requested_positive + CHECK (num_parts_requested >= 1), + + CONSTRAINT chk_model_upload_session_file_size_bytes_positive + CHECK (file_size_bytes > 0), + + CONSTRAINT chk_model_upload_session_part_size_bytes_positive + CHECK (part_size_bytes > 0), + + CONSTRAINT chk_model_upload_session_part_size_bytes_s3_upper_bound + CHECK (part_size_bytes <= 5368709120) +); + +-- model_upload_session_part +CREATE TABLE IF NOT EXISTS model_upload_session_part +( + upload_id VARCHAR(256) NOT NULL, + part_number INT NOT NULL, + etag TEXT NOT NULL DEFAULT '', + + PRIMARY KEY (upload_id, part_number), + + CONSTRAINT chk_model_part_number_positive CHECK (part_number > 0), + + FOREIGN KEY (upload_id) + REFERENCES model_upload_session(upload_id) + ON DELETE CASCADE +); + -- operator_executions (modified to match MySQL: no separate primary key; added console_messages_uri) CREATE TABLE IF NOT EXISTS operator_executions ( diff --git a/sql/updates/41.sql b/sql/updates/41.sql new file mode 100644 index 0000000000..8a3590b2b6 --- /dev/null +++ b/sql/updates/41.sql @@ -0,0 +1,82 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +\c texera_db + +SET search_path TO texera_db; + +BEGIN; + +-- Session-based multipart upload for model files. Tracks in-progress multipart +-- uploads so a model version can be assembled from parts and resumed across requests. + +CREATE TABLE IF NOT EXISTS model_upload_session +( + mid INT NOT NULL, + uid INT NOT NULL, + file_path TEXT NOT NULL, + upload_id VARCHAR(256) NOT NULL UNIQUE, + physical_address TEXT, + num_parts_requested INT NOT NULL, + file_size_bytes BIGINT NOT NULL, + part_size_bytes BIGINT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + + PRIMARY KEY (uid, mid, file_path), + + FOREIGN KEY (mid) REFERENCES model(mid) ON DELETE CASCADE, + FOREIGN KEY (uid) REFERENCES "user"(uid) ON DELETE CASCADE, + + CONSTRAINT chk_model_upload_session_num_parts_requested_positive + CHECK (num_parts_requested >= 1), + + CONSTRAINT chk_model_upload_session_file_size_bytes_positive + CHECK (file_size_bytes > 0), + + CONSTRAINT chk_model_upload_session_part_size_bytes_positive + CHECK (part_size_bytes > 0), + + CONSTRAINT chk_model_upload_session_part_size_bytes_s3_upper_bound + CHECK (part_size_bytes <= 5368709120) +); + +CREATE TABLE IF NOT EXISTS model_upload_session_part +( + upload_id VARCHAR(256) NOT NULL, + part_number INT NOT NULL, + etag TEXT NOT NULL DEFAULT '', + + PRIMARY KEY (upload_id, part_number), + + CONSTRAINT chk_model_part_number_positive CHECK (part_number > 0), + + FOREIGN KEY (upload_id) + REFERENCES model_upload_session(upload_id) + ON DELETE CASCADE +); + +-- Version names are generated from an unlocked count(*), and FileResolver looks a version +-- up by (mid, name) with fetchOneInto. Enforce the uniqueness the generator assumes. +ALTER TABLE model_version + DROP CONSTRAINT IF EXISTS uq_model_version_mid_name; + +ALTER TABLE model_version + ADD CONSTRAINT uq_model_version_mid_name UNIQUE (mid, name); + +COMMIT;
