Skip to content
Open
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
50 changes: 35 additions & 15 deletions src/Lean/Attributes.lean
Original file line number Diff line number Diff line change
Expand Up @@ -260,22 +260,29 @@ structure ParametricAttributeImpl (α : Type) extends AttributeImplCore where
filterExport : Environment → Name → α → Bool := fun env n _ =>
env.contains (skipRealize := false) n

def registerParametricAttribute (impl : ParametricAttributeImpl α) : IO (ParametricAttribute α) := do
let ext : PersistentEnvExtension (Name × α) (Name × α) (List Name × NameMap α) ← registerPersistentEnvExtension {
name := impl.ref
def registerParametricAttributeExt (ref : Name) (preserveOrder : Bool := false)
(filterExport : Environment → Name → α → Bool := fun env n _ =>
env.contains (skipRealize := false) n) :
IO (PersistentEnvExtension (Name × α) (Name × α) (List Name × NameMap α)) :=
registerPersistentEnvExtension {
name := ref
mkInitial := pure ([], {})
addImportedFn := fun _ => pure ([], {})
addEntryFn := fun (decls, m) (p : Name × α) => (p.1 :: decls, m.insert p.1 p.2)
exportEntriesFnEx := fun env (decls, m) => Id.run do
let all := if impl.preserveOrder then
let all := if preserveOrder then
decls.toArray.reverse.filterMap (fun n => return (n, ← m.find? n))
else
let r := m.foldl (fun a n p => a.push (n, p)) #[]
r.qsort (fun a b => Name.quickLt a.1 b.1)
let exported := all.filter (fun ⟨n, a⟩ => impl.filterExport env n a)
let exported := all.filter (fun ⟨n, a⟩ => filterExport env n a)
{ exported, server := exported, «private» := all }
statsFn := fun (_, m) => "parametric attribute" ++ Format.line ++ "number of local entries: " ++ format m.size
}

def registerParametricAttributeForExt (impl : ParametricAttributeImpl α)
(ext : PersistentEnvExtension (Name × α) (Name × α) (List Name × NameMap α)) :
IO (ParametricAttribute α) := do
let attrImpl : AttributeImpl := {
impl.toAttributeImplCore with
add := fun decl stx kind => do
Expand All @@ -290,27 +297,40 @@ def registerParametricAttribute (impl : ParametricAttributeImpl α) : IO (Parame
registerBuiltinAttribute attrImpl
pure { attr := attrImpl, ext, preserveOrder := impl.preserveOrder }

def registerParametricAttribute (impl : ParametricAttributeImpl α) : IO (ParametricAttribute α) := do
let ext ← registerParametricAttributeExt (α := α) impl.ref impl.preserveOrder impl.filterExport
registerParametricAttributeForExt impl ext

namespace ParametricAttribute

def getParam? [Inhabited α] (attr : ParametricAttribute α) (env : Environment) (decl : Name) : Option α :=
def getParamFromExt? [Inhabited α]
(ext : PersistentEnvExtension (Name × α) (Name × α) (List Name × NameMap α))
(preserveOrder : Bool) (env : Environment) (decl : Name) : Option α :=
match env.getModuleIdxFor? decl with
| some modIdx =>
let entry? := if attr.preserveOrder then
(attr.ext.getModuleEntries env modIdx).find? (·.1 == decl)
let entry? := if preserveOrder then
(ext.getModuleEntries env modIdx).find? (·.1 == decl)
else
(attr.ext.getModuleEntries env modIdx).binSearch (decl, default) (fun a b => Name.quickLt a.1 b.1)
(ext.getModuleEntries env modIdx).binSearch (decl, default) (fun a b => Name.quickLt a.1 b.1)
match entry? with
| some (_, val) => some val
| none => none
| none => (attr.ext.getState env).2.find? decl
| none => (ext.getState env).2.find? decl

def setParam (attr : ParametricAttribute α) (env : Environment) (decl : Name) (param : α) : Except String Environment :=
def getParam? [Inhabited α] (attr : ParametricAttribute α) (env : Environment) (decl : Name) : Option α :=
getParamFromExt? attr.ext attr.preserveOrder env decl

def setParamFromExt
(ext : PersistentEnvExtension (Name × α) (Name × α) (List Name × NameMap α)) (attr : AttributeImpl) (env : Environment) (decl : Name) (param : α) : Except String Environment :=
if (env.getModuleIdxFor? decl).isSome then
Except.error (s!"Failed to add parametric attribute `[{attr.attr.name}]` to `{decl}`: Declaration is in an imported module")
else if ((attr.ext.getState env).2.find? decl).isSome then
Except.error (s!"Failed to add parametric attribute `[{attr.attr.name}]` to `{decl}`: Attribute has already been set")
Except.error (s!"Failed to add parametric attribute `[{attr.name}]` to `{decl}`: Declaration is in an imported module")
else if ((ext.getState env).2.find? decl).isSome then
Except.error (s!"Failed to add parametric attribute `[{attr.name}]` to `{decl}`: Attribute has already been set")
else
Except.ok (attr.ext.addEntry env (decl, param))
Except.ok (ext.addEntry env (decl, param))

def setParam (attr : ParametricAttribute α) (env : Environment) (decl : Name) (param : α) : Except String Environment :=
setParamFromExt attr.ext attr.attr env decl param

end ParametricAttribute

Expand Down
Loading