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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 24 additions & 1 deletion Sources/FluidAudio/Shared/Download/HFTreeLister.swift
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,22 @@ import Foundation

/// A file discovered in a HuggingFace repository tree listing.
struct RemoteFile: Equatable, Sendable {
enum ContentID: Equatable, Sendable {
case lfsSHA256(String)
case gitBlobSHA1(String)
}

let path: String
/// Size in bytes as reported by the tree API; `-1` when not reported.
let size: Int
/// LFS content hash, or the Git blob hash for a regular file (not an LFS pointer).
let contentID: ContentID?

init(path: String, size: Int, contentID: ContentID? = nil) {
self.path = path
self.size = size
self.contentID = contentID
}
}

/// The one tree-listing implementation for the download stack (#765 Wave 3),
Expand Down Expand Up @@ -93,7 +106,17 @@ enum HFTreeLister {
)
} else if itemType == "file" {
guard include(itemPath, false) else { continue }
files.append(RemoteFile(path: itemPath, size: item["size"] as? Int ?? -1))
let contentID: RemoteFile.ContentID?
if item["lfs"] != nil {
// The top-level oid of an LFS entry hashes its pointer,
// not the resolved content downloaded into the cache.
let lfs = item["lfs"] as? [String: Any]
contentID = (lfs?["oid"] as? String).map { .lfsSHA256($0) }
} else {
contentID = (item["oid"] as? String).map { .gitBlobSHA1($0) }
}
files.append(
RemoteFile(path: itemPath, size: item["size"] as? Int ?? -1, contentID: contentID))
}
}

Expand Down
176 changes: 176 additions & 0 deletions Sources/FluidAudio/Shared/Download/ModelCache.swift
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import CryptoKit
import Foundation

/// On-disk model-cache knowledge for the download stack (#765 Wave 4):
Expand Down Expand Up @@ -29,6 +30,181 @@ enum ModelCache {
return storedRevision == revision
}

/// Inventory every cached compiled bundle and plain file before listing a
/// legacy pinned cache. A revision marker covers all variants in this folder.
static func legacyCacheContents(
at repoPath: URL, revision: String
) throws -> (bundles: Set<String>, files: Set<String>)? {
let fm = FileManager.default
var isDirectory: ObjCBool = false
guard revision != "main",
!fm.fileExists(atPath: repoPath.appendingPathComponent(revisionMarkerName).path),
fm.fileExists(atPath: repoPath.path, isDirectory: &isDirectory), isDirectory.boolValue
else { return nil }

let root = repoPath.resolvingSymlinksInPath()
var bundles: Set<String> = root.pathExtension == "mlmodelc" ? [""] : []
var files: Set<String> = []
guard let enumerator = fm.enumerator(at: root, includingPropertiesForKeys: nil) else {
return (bundles, files)
}
for case let item as URL in enumerator {
let item = item.resolvingSymlinksInPath()
let relative = String(item.path.dropFirst(root.path.count + 1))
if item.pathExtension == "mlmodelc" {
bundles.insert(relative)
enumerator.skipDescendants()
continue
}
guard !relative.hasSuffix(".partial"), !relative.hasSuffix(".partial.etag"),
let attributes = try attributesIfPresent(at: item),
attributes[.type] as? FileAttributeType == .typeRegular
else { continue }
files.insert(relative)
}
return (bundles, files)
}

/// Adopt only after the pinned listing covers every locally cached bundle.
/// Incomplete bundles are removed as a unit so the model-existence gate
/// re-enters download after an interruption. Existing markers remain authoritative.
static func adoptLegacyCache(
at repoPath: URL, revision: String, files: [RemoteFile], subPath: String? = nil
) throws {
let fm = FileManager.default
guard let contents = try legacyCacheContents(at: repoPath, revision: revision) else { return }
let marker = repoPath.appendingPathComponent(revisionMarkerName)
var listedBundles: Set<String> = []
var listedFiles: Set<String> = []
var invalidBundles: Set<String> = []
var invalidFiles: Set<URL> = []
var keptFiles: [(url: URL, bundle: String?)] = []
for file in files {
let localPath = localPath(for: file.path, subPath: subPath)
let components = localPath.split(separator: "/")
let bundle =
repoPath.pathExtension == "mlmodelc"
? ""
: components.firstIndex(where: { $0.hasSuffix(".mlmodelc") }).map {
components[...$0].joined(separator: "/")
}
if let bundle { listedBundles.insert(bundle) }
listedFiles.insert(localPath)
let destination = repoPath.appendingPathComponent(localPath)
guard try hasMatchingContent(file, at: destination) else {
if let bundle {
invalidBundles.insert(bundle)
} else {
invalidFiles.insert(destination)
}
continue
}
keptFiles.append((destination, bundle))
}
invalidBundles.formUnion(contents.bundles.subtracting(listedBundles))
for file in contents.files.subtracting(listedFiles) {
invalidFiles.insert(repoPath.appendingPathComponent(file))
}
var removals = invalidBundles.sorted().map {
$0.isEmpty ? repoPath : repoPath.appendingPathComponent($0)
}
removals.append(contentsOf: invalidFiles)
// Rejected or missing files must not return through finished-partial reuse.
for file in invalidFiles {
removals.append(file.appendingPathExtension("partial"))
removals.append(file.appendingPathExtension("partial.etag"))
}
for file in keptFiles where file.bundle.map({ !invalidBundles.contains($0) }) ?? true {
removals.append(file.url.appendingPathExtension("partial"))
removals.append(file.url.appendingPathExtension("partial.etag"))
}
for path in removals {
do {
try fm.removeItem(at: path)
} catch {
guard isMissingFile(error) else { throw error }
}
}
try fm.createDirectory(at: repoPath, withIntermediateDirectories: true)
try Data((revision + "\n").utf8).write(to: marker, options: .atomic)
}

/// Size is only a pre-check: an unknown size can still match its content ID.
private static func hasMatchingContent(_ file: RemoteFile, at destination: URL) throws -> Bool {
guard let contentID = file.contentID,
let attributes = try attributesIfPresent(at: destination),
attributes[.type] as? FileAttributeType == .typeRegular,
let size = (attributes[.size] as? NSNumber)?.int64Value,
file.size < 0 || size == Int64(file.size)
else { return false }

let expected: String
let hexLength: Int
switch contentID {
case .lfsSHA256(let oid):
expected = oid.lowercased()
hexLength = 64
case .gitBlobSHA1(let oid):
expected = oid.lowercased()
hexLength = 40
}
guard expected.utf8.count == hexLength,
expected.utf8.allSatisfy({ (48...57).contains($0) || (97...102).contains($0) })
else { return false }

do {
switch contentID {
case .lfsSHA256:
return try hashFile(at: destination, using: SHA256()) == expected
case .gitBlobSHA1:
var hasher = Insecure.SHA1()
hasher.update(data: Data("blob \(size)\0".utf8))
return try hashFile(at: destination, using: hasher) == expected
}
} catch {
guard isMissingFile(error) else { throw error }
return false
}
}

/// Stream large weights in bounded chunks rather than materializing them in memory.
private static func hashFile<H: HashFunction>(at url: URL, using initialHasher: H) throws -> String {
let handle = try FileHandle(forReadingFrom: url)
defer { try? handle.close() }
var hasher = initialHasher
while let chunk = try handle.read(upToCount: 1_048_576), !chunk.isEmpty {
hasher.update(data: chunk)
}
return hasher.finalize().map { String(format: "%02x", $0) }.joined()
}

static func localPath(for remotePath: String, subPath: String?) -> String {
guard let subPath, remotePath.hasPrefix("\(subPath)/") else { return remotePath }
return String(remotePath.dropFirst(subPath.count + 1))
}

/// Missing-file races are harmless; permission and other I/O errors still propagate.
static func isMissingFile(_ error: Error) -> Bool {
let error = error as NSError
if error.domain == NSCocoaErrorDomain,
error.code == NSFileNoSuchFileError || error.code == NSFileReadNoSuchFileError
{
return true
}
if error.domain == NSPOSIXErrorDomain, error.code == Int(POSIXErrorCode.ENOENT.rawValue) { return true }
guard let underlying = error.userInfo[NSUnderlyingErrorKey] as? Error else { return false }
return isMissingFile(underlying)
}

private static func attributesIfPresent(at url: URL) throws -> [FileAttributeKey: Any]? {
do {
return try FileManager.default.attributesOfItem(atPath: url.path)
} catch {
guard isMissingFile(error) else { throw error }
return nil
}
}

/// Prepare a managed cache for downloads from one resolved revision.
/// Existing files are preserved when the marker matches and replaced when
/// the requested revision changes. The marker is written before downloads
Expand Down
53 changes: 52 additions & 1 deletion Sources/FluidAudio/Shared/Download/ModelHub.swift
Original file line number Diff line number Diff line change
Expand Up @@ -502,7 +502,16 @@ public enum ModelHub {
}
}

let requestedPaths = Set(filesToDownload.map(\.path))
// Validate all cached variants before stamping the shared marker,
// but only fetch files selected for this caller's download.
filesToDownload = try await filesForLegacyAdoption(
filesToDownload, at: repoPath, repo: repo, revision: revision, subPath: subPath,
includeRepoRootFiles: true, fetch: treeFetch)
try ModelCache.adoptLegacyCache(
at: repoPath, revision: revision, files: filesToDownload, subPath: subPath)
try ModelCache.prepareForDownload(at: repoPath, revision: revision)
filesToDownload.removeAll { !requestedPaths.contains($0.path) }
logger.info("Found \(filesToDownload.count) files to download")

// Compute total known bytes for byte-weighted progress.
Expand Down Expand Up @@ -613,15 +622,22 @@ public enum ModelHub {
let reporter = ProgressReporter(handler: progressHandler, downloadPhaseWeight: 1.0)
reporter.listing()
let revision = ModelRegistry.mapRevision(repo.remotePath, default: repo.revision)
let filesToDownload: [RemoteFile] = try await HFTreeLister.listTree(
var filesToDownload: [RemoteFile] = try await HFTreeLister.listTree(
repoRemotePath: repo.remotePath,
revision: revision,
startingAt: subdirectory,
include: { itemPath, _ in shouldSkip?(itemPath) != true },
fetch: HFTreeLister.fetch(using: listingSession)
)
let revisionCache = repoDirectory.appendingPathComponent(subdirectory)
let requestedPaths = Set(filesToDownload.map(\.path))
filesToDownload = try await filesForLegacyAdoption(
filesToDownload, at: revisionCache, repo: repo, revision: revision, subPath: subdirectory,
includeRepoRootFiles: false, fetch: HFTreeLister.fetch(using: listingSession))
try ModelCache.adoptLegacyCache(
at: revisionCache, revision: revision, files: filesToDownload, subPath: subdirectory)
try ModelCache.prepareForDownload(at: revisionCache, revision: revision)
filesToDownload.removeAll { !requestedPaths.contains($0.path) }
let totalFiles = filesToDownload.count
logger.info("Found \(totalFiles) files in \(subdirectory)")

Expand Down Expand Up @@ -669,6 +685,41 @@ public enum ModelHub {
logger.info("Downloaded \(subdirectory) from \(repo.folderName)")
}

/// A legacy marker covers every cached variant, so extend the calling
/// variant's listing to all local compiled bundles and plain files first.
private static func filesForLegacyAdoption(
_ files: [RemoteFile], at repoPath: URL, repo: Repo, revision: String, subPath: String?,
includeRepoRootFiles: Bool, fetch: HFTreeLister.Fetch
) async throws -> [RemoteFile] {
guard let contents = try ModelCache.legacyCacheContents(at: repoPath, revision: revision) else { return files }
var additional = try await HFTreeLister.listTree(
repoRemotePath: repo.remotePath, revision: revision, startingAt: subPath ?? "",
include: { path, isDirectory in
let local = ModelCache.localPath(for: path, subPath: subPath)
let inBundle = contents.bundles.contains {
$0.isEmpty || local == $0 || local.hasPrefix($0 + "/")
|| (isDirectory && $0.hasPrefix(local + "/"))
}
let plainFile =
contents.files.contains(local)
|| (isDirectory && contents.files.contains { $0.hasPrefix(local + "/") })
return inBundle || plainFile
}, fetch: fetch)
let covered = Set((files + additional).map { ModelCache.localPath(for: $0.path, subPath: subPath) })
let missingRoots = Set(contents.files.filter { !$0.contains("/") }).subtracting(covered)
// Only repo downloads flatten a subPath and can share root auxiliaries.
// Subdirectory downloads preserve remote paths and never use root files.
if includeRepoRootFiles, subPath != nil, !missingRoots.isEmpty {
additional += try await HFTreeLister.listTree(
repoRemotePath: repo.remotePath, revision: revision,
include: { path, isDirectory in !isDirectory && missingRoots.contains(path) }, fetch: fetch)
}
var merged = files
var paths = Set(files.map(\.path))
for file in additional where paths.insert(file.path).inserted { merged.append(file) }
return merged
}

/// One file of a subdirectory download, run inside the bounded task group.
private static func downloadSubdirectoryFile(
_ file: RemoteFile,
Expand Down
28 changes: 28 additions & 0 deletions Tests/FluidAudioTests/Shared/HFTreeListerTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,34 @@ final class HFTreeListerTests: XCTestCase {

// MARK: - Walking + pruning

func testParsesContentIDsWithoutFallingBackToLFSPointerOID() async throws {
let server = PageServer()
let sha256 = String(repeating: "a", count: 64)
let blobSHA1 = "e69de29bb2d1d6434b8b29ae775ad8c2e48c5391"
try server.addPage(
url: treeURL(),
items: [
["path": "weight.bin", "type": "file", "size": 5, "oid": blobSHA1, "lfs": ["oid": sha256]],
["path": "empty.json", "type": "file", "size": 0, "oid": blobSHA1],
["path": "missing-lfs.bin", "type": "file", "oid": blobSHA1, "lfs": ["size": 5]],
["path": "malformed-lfs.bin", "type": "file", "oid": blobSHA1, "lfs": "invalid"],
["path": "missing-id.json", "type": "file"],
])

let files = try await HFTreeLister.listTree(
repoRemotePath: Self.repo, include: { _, _ in true }, fetch: server.fetch)

XCTAssertEqual(
files,
[
RemoteFile(path: "weight.bin", size: 5, contentID: .lfsSHA256(sha256)),
RemoteFile(path: "empty.json", size: 0, contentID: .gitBlobSHA1(blobSHA1)),
RemoteFile(path: "missing-lfs.bin", size: -1),
RemoteFile(path: "malformed-lfs.bin", size: -1),
RemoteFile(path: "missing-id.json", size: -1),
])
}

func testRecursiveWalkWithPruningAndFileExclusion() async throws {
let server = PageServer()
try server.addPage(
Expand Down
Loading
Loading