summaryrefslogtreecommitdiffstats
path: root/pkgs/build-support/fetchhuggingface/default.nix
blob: 9bf33b2486554265c51bed72b0ed871ca0921687 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
{
  lib,
  repoRevToNameMaybe,
  fetchgit,
}:

let
  repoPrefixes = {
    model = "";
    dataset = "datasets/";
    space = "spaces/";
  };
in

lib.makeOverridable (
  {
    repoId,
    tag ? null,
    rev ? null,
    name ? repoRevToNameMaybe repoId (lib.revOrTag rev tag) "huggingface",
    domain ? "huggingface.co",
    repoType ? "model",
    backend ? "xet",
    branchName ? null,
    deepClone ? false,
    fetchSubmodules ? false,
    fetchTags ? false,
    leaveDotGit ? null,
    rootDir ? "",
    sparseCheckout ? null,
    passthru ? { },
    meta ? { },
    ... # For hash agility and additional fetchgit arguments
  }@args:

  assert (
    lib.assertMsg (lib.xor (tag == null) (
      rev == null
    )) "fetchFromHuggingFace requires one of either `rev` or `tag` to be provided (not both)."
  );

  assert (lib.assertOneOf "repoType" repoType (builtins.attrNames repoPrefixes));
  assert (
    lib.assertOneOf "backend" backend [
      "lfs"
      "xet"
    ]
  );

  let
    position = (
      if args.meta.description or null != null then
        builtins.unsafeGetAttrPos "description" args.meta
      else if tag != null then
        builtins.unsafeGetAttrPos "tag" args
      else
        builtins.unsafeGetAttrPos "rev" args
    );
    baseUrl = "https://${domain}/${repoPrefixes.${repoType}}${repoId}";
    gitRepoUrl = "${baseUrl}.git";
    newMeta =
      meta
      // {
        homepage = meta.homepage or baseUrl;
      }
      // lib.optionalAttrs (position != null) {
        # to indicate where derivation originates, similar to make-derivation.nix's mkDerivation
        position = "${position.file}:${toString position.line}";
      };
    backendFetcher = builtins.getAttr backend {
      lfs = fetchgit;
      xet = throw "fetchFromHuggingFace: the Xet backend is not implemented yet";
    };
  in
  assert (
    lib.assertMsg (
      builtins.match "[^/]+(/[^/]+)?" repoId != null
    ) "fetchFromHuggingFace requires `repoId` to be in the form `repo` or `owner/repo`."
  );
  backendFetcher (
    removeAttrs args [
      "backend"
      "domain"
      "repoId"
      "repoType"
    ]
    // {
      inherit
        branchName
        deepClone
        fetchSubmodules
        fetchTags
        leaveDotGit
        name
        rootDir
        sparseCheckout
        tag
        rev
        ;
      url = gitRepoUrl;
      fetchLFS = true;
      meta = newMeta;
      passthru = {
        inherit gitRepoUrl;
      }
      // passthru;
    }
  )
  // {
    inherit
      repoId
      repoType
      ;
  }
)