This is an automated email from the ASF dual-hosted git repository.
github-merge-queue[bot] pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/texera.git
The following commit(s) were added to refs/heads/main by this push:
new 7fbf64e3da feat(file-service): add model file upload and version
endpoints (#6872)
7fbf64e3da is described below
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;