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
14 changes: 7 additions & 7 deletions hkmc2/shared/src/main/scala/hkmc2/AsyncLowering.scala
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ class AsyncLowering(using TL, Raise, Elaborator.State, Elaborator.Ctx, Config):
pl.allParams.map: p =>
val v = p.sym
val nv = VarSymbol(v.id, erasedType = v.erasedType)
(p, p.copy(sym = nv))
(p, p.copy(sym = nv)(p.toLoc))
val symMap = outerParams.iterator.map(p => p._1.sym -> p._2.sym).toMap[SimpleSymbol, SimpleSymbol]
val thisVar = VarSymbol(Tree.Ident("this"), erasedType = fun.owner.flatMap(_.asThis.erasedValueType))
val thisParam = fun.owner.map(_ => Param.simple(thisVar))
Expand All @@ -56,22 +56,22 @@ class AsyncLowering(using TL, Raise, Elaborator.State, Elaborator.Ctx, Config):
)
)
val vars = fun.params.flatMap(_.paramSyms)
val noAsync = fun.annotations.filterNot(_ is Annot.Async)
val noAsync = fun.annotations.filterNot(_.isInstanceOf[Annot.Async])
val transformer = new BlockTransformer(SymbolSubst.Id):
override def applySimpleSymbol(sym: SimpleSymbol): SimpleSymbol =
symMap.getOrElse(sym, sym)
override def applyValue(v: Value)(k: Value => Block): Block = v match
case Value.This(sym) if fun.owner.contains(sym) =>
k(Value.SimpleRef(thisVar))
k(Value.SimpleRef(thisVar)(v.toLoc))
case _ => super.applyValue(v)(k)
val newBody = transformer.applyBlock(wrapAwait(true)(applyFunBodyLikeBlock(fun.body)))
collectedFunDefn += FunDefn(N, outerBms, outerDsym, PlainParamList((thisParam.iterator ++ outerParams.iterator.map(_._2)).toList) :: PlainParamList(Nil) :: Nil, newBody)(fun.configOverride, noAsync)
val callArgs = (fun.owner.iterator.map(s => Arg(N, Value.This(s))) ++ fun.params.iterator.flatMap(_.allParams.iterator.map(p => Arg(N, Value.SimpleRef(p.sym))))).toList
val outerCall = Call(Value.MemberRef(outerBms, outerDsym), callArgs ne_:: Nil)(CallMetadata.mlsFunWithEffect)
collectedFunDefn += FunDefn(N, outerBms, outerDsym, PlainParamList((thisParam.iterator ++ outerParams.iterator.map(_._2)).toList)(N) :: PlainParamList(Nil)(N) :: Nil, newBody)(fun.configOverride, noAsync)
val callArgs = (fun.owner.iterator.map(s => Arg(N, Value.This(s)(N))) ++ fun.params.iterator.flatMap(_.allParams.iterator.map(p => Arg(N, Value.SimpleRef(p.sym)(N))))).toList
val outerCall = Call(Value.MemberRef(outerBms, outerDsym)(N), callArgs ne_:: Nil)(CallMetadata.mlsFunWithEffect, N)
val tmp = TempSymbol(N, erasedType = outerCall.erasedValueType, "tmp")
val wrapperBody = blockBuilder
.assignScoped(tmp, outerCall)
.ret(Call(Value.SimpleRef(State.runtimeSymbol).selSN("toJsAsync"), (tmp.asSimpleRef.asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun))
.ret(Call(Value.SimpleRef(State.runtimeSymbol)(N).selSN("toJsAsync"), (tmp.asSimpleRef.asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun, N))
FunDefn(fun.owner, fun.sym, fun.dSym, fun.params, wrapperBody)(fun.configOverride, noAsync)

override def applyMainBlock(main: Block): Block =
Expand Down
4 changes: 2 additions & 2 deletions hkmc2/shared/src/main/scala/hkmc2/CompilerCtx.scala
Original file line number Diff line number Diff line change
Expand Up @@ -121,10 +121,10 @@ class CompilerCtx(
case _ => t.subTerms.exists(findQuote)
val hasQuote = findQuote(blk0)
val blk = new Term.Blk(
Import(State.runtimeSymbol, paths.runtimeFile.toString, paths.runtimeFile) ::
Import(State.runtimeSymbol, paths.runtimeFile.toString, paths.runtimeFile)(N) ::
// Only import `Term.mls` when necessary.
(if hasQuote then
Import(State.termSymbol, paths.termFile.toString, paths.termFile) :: blk0.stats
Import(State.termSymbol, paths.termFile.toString, paths.termFile)(N) :: blk0.stats
else
blk0.stats),
blk0.res
Expand Down
162 changes: 84 additions & 78 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/Block.scala

Large diffs are not rendered by default.

16 changes: 8 additions & 8 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/BlockSimplifier.scala
Original file line number Diff line number Diff line change
Expand Up @@ -289,7 +289,7 @@ class BlockSimplifier
case Value.SimpleRef(loc: LocalVarSymbol) if localVars.contains(loc) && !definedVars.contains(loc) =>
registerChange(s"${loc.showDbg} is never assigned; replacing read with undefined")
// if !symbolsToPreserve(loc) then removedLocals += loc
k(Value.Lit(syntax.Tree.UnitLit(false)))
k(Value.Lit(syntax.Tree.UnitLit(false))(v.toLoc))
case _ => super.applyValue(v)(k)

override def applyBlock(b: Block): Block = b match
Expand Down Expand Up @@ -863,7 +863,7 @@ class BlockSimplifier
registerChange(s"immediate assigned call prefix ${lhs.showDbg} ~> ${path.showDbg}")
applyPath(path): path2 =>
val lhs2 = recordAssignmentFact(lhs, path2, ass)
val combined = Call(path2, argss)(call.metadata).withLocOf(call)
val combined = Call(path2, argss)(call.metadata, call.toLoc)
val res = applyBlock(Assign(nextLhs, combined, rst))
// * Note that it is incorrect to eliminate the `lhs` assignment even if `!rst.freeVars(lhs)`,
// * because the assignment may be visible from an outer block
Expand All @@ -878,7 +878,7 @@ class BlockSimplifier
registerChange(s"immediate returned call prefix ${lhs.showDbg} ~> ${path.showDbg}")
applyPath(path): path2 =>
val lhs2 = recordAssignmentFact(lhs, path2, ass)
val combined = Call(path2, argss)(call.metadata).withLocOf(call)
val combined = Call(path2, argss)(call.metadata, call.toLoc)
val res = applyBlock(Return(combined))
if symbolsToPreserve(lhs) then Assign(lhs2, path2, res) else res

Expand Down Expand Up @@ -1157,7 +1157,7 @@ class BlockSimplifier
analysis.litValue match
case true =>
registerChange(s"${loc.showDbg} ~> undefined")
return k(Value.Lit(syntax.Tree.UnitLit(false)))
return k(Value.Lit(syntax.Tree.UnitLit(false))(v.toLoc))
case lit: (Value | Cast) =>
registerChange(s"${loc.showDbg} ~> ${lit.showDbg}")
return k(lit)
Expand Down Expand Up @@ -1203,7 +1203,7 @@ class BlockSimplifier
prefix.metadata.mayRaiseEffects || c.metadata.mayRaiseEffects,
prefix.metadata.annotations ++ c.metadata.annotations,
),
).withLocOf(c)
c.toLoc)
super.applyResult(combined)(k)
case N => super.applyResult(r)(k)

Expand Down Expand Up @@ -1474,7 +1474,7 @@ class BlockSimplifier
if args.size < params.params.size then return N
val (fixedArgs, restArgs) = args.splitAt(params.params.size)
S(fixedArgs.zip(params.params).map((arg, param) => (param.sym, arg.value)) ++
List((params.restParam.get.sym, Tuple(true, restArgs))))
List((params.restParam.get.sym, Tuple(true, restArgs)(N))))

/** Match multiple argument lists against multiple parameter lists.
* Returns None if any arg list fails to match its corresponding param list,
Expand Down Expand Up @@ -1917,8 +1917,8 @@ class BlockSimplifier
acc(Scoped(Set(resSym), newBlk(
k(Call(resSym.asSimpleRef, extraArgss.ne_!)(
call.metadata.copy(
annotations = call.metadata.annotations.filterNot(_ == Annot.TailCall),
))))))
annotations = call.metadata.annotations.filterNot(_.isInstanceOf[Annot.TailCall]),
), call.toLoc)))))
case (sym, value) :: argRest =>
val newSym = VarSymbol(sym.id, erasedType = sym.erasedType)
go(acc.assignScoped(newSym, value), argRest, mapping + (sym -> newSym))
Expand Down
20 changes: 10 additions & 10 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/BlockTransformer.scala
Original file line number Diff line number Diff line change
Expand Up @@ -164,33 +164,33 @@ class BlockTransformer(subst: SymbolSubst):
applyPath(fun): fun2 =>
applyListOf(argss, (args, k2) => applyArgs(args)(k2)): argss2 =>
k(if (fun2 is fun) && (argss2 is argss) then r
else Call(fun2, argss2.ne_!)(r.metadata).withLocOf(r))
else Call(fun2, argss2.ne_!)(r.metadata, r.toLoc))
case r @ Instantiate(mut, cls, argss) =>
applyPath(cls): cls2 =>
applyListOf(argss, (args, k2) => applyArgs(args)(k2)): argss2 =>
k(if (cls2 is cls) && (argss2 is argss) then r
else Instantiate(mut, cls2, argss2)(r.metadata).withLocOf(r))
else Instantiate(mut, cls2, argss2)(r.metadata, r.toLoc))
case l: Lambda => k(applyLam(l))
case Tuple(mut, elems) =>
applyArgs(elems): elems2 =>
k(if (elems2 is elems) then r else Tuple(mut, elems2).withLocOf(r))
k(if (elems2 is elems) then r else Tuple(mut, elems2)(r.toLoc))
case Record(mut, fields) =>
applyRcdArgs(fields): fields2 =>
k(if fields2 is fields then r else Record(mut, fields2).withLocOf(r))
k(if fields2 is fields then r else Record(mut, fields2)(r.toLoc))
case p: Path => applyPath(p)(k)

def applyPath(p: Path)(k: Path => Block): Block = p match
case DynSelect(qual, fld, arrayIdx) =>
applyPath(qual): qual2 =>
applyPath(fld): fld2 =>
k(if (qual2 is qual) && (fld2 is fld) then p else DynSelect(qual2, fld2, arrayIdx).withLocOf(p))
k(if (qual2 is qual) && (fld2 is fld) then p else DynSelect(qual2, fld2, arrayIdx)(p.toLoc))
case p @ Select(qual, name) =>
applyPath(qual): qual2 =>
val sym2 = p.symbol.mapConserve(_.subst)
k(if (qual2 is qual) && (sym2 is p.symbol) then p else Select(qual2, name)(sym2)(p.sanitize).withLocOf(p))
k(if (qual2 is qual) && (sym2 is p.symbol) then p else Select(qual2, name)(sym2, p.toLoc)(p.sanitize))
case c @ Cast(value, target, check) =>
applyResult(value): value2 =>
k(if value2 is value then c else Cast(value2, target, check).withLocOf(c))
k(if value2 is value then c else Cast(value2, target, check)(c.toLoc))
case v: Value => applyValue(v)(k)

def applyValue(v: Value)(k: Value => Block) = v match
Expand Down Expand Up @@ -300,11 +300,11 @@ class BlockTransformer(subst: SymbolSubst):
def applyParamList(pl: ParamList): ParamList =
def applyParam(p: Param): Param =
val sym2 = p.sym.subst
if sym2 is p.sym then p else p.copy(sym = sym2)
if sym2 is p.sym then p else p.copy(sym = sym2)(p.toLoc)
val params2 = pl.params.mapConserve(applyParam)
val rest2 = pl.restParam.mapConserve(applyParam)
if (params2 is pl.params) && (rest2 is pl.restParam)
then pl else ParamList(pl.flags, params2, rest2)
then pl else ParamList(pl.flags, params2, rest2)(pl.toLoc)

def applyCase(cse: Case)(k: Case => Block): Block = cse match
case Case.Lit(lit) => k(cse)
Expand All @@ -327,7 +327,7 @@ class BlockTransformer(subst: SymbolSubst):
def applyLam(lam: Lambda): Lambda =
val params2 = applyParamList(lam.params)
val body2 = applyFunBodyLikeBlock(lam.body)
if (params2 is lam.params) && (body2 is lam.body) then lam else Lambda(params2, body2)(lam.annot)
if (params2 is lam.params) && (body2 is lam.body) then lam else Lambda(params2, body2)(lam.annot, lam.toLoc)

def applyListOf[A](ls: List[A], f: (A, (A => Block)) => Block)(k: List[A] => Block): Block =
def rec(ls: List[A], k: List[A] => Block): Block = ls match
Expand Down
20 changes: 10 additions & 10 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/BufferableTransform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -34,16 +34,16 @@ class BufferableTransform()(using State, Raise):
(sym, VarSymbol(sym.id, erasedType = N))
.toMap
def mapParam(p: Param) =
Param(p.flags, varMap(p.sym), p.sign, p.modulefulness)
(params.map(pl => ParamList(pl.flags, pl.params.map(mapParam), pl.restParam.map(mapParam))), varMap.toMap)
Param(p.flags, varMap(p.sym), p.sign, p.modulefulness)(p.toLoc)
(params.map(pl => ParamList(pl.flags, pl.params.map(mapParam), pl.restParam.map(mapParam))(pl.toLoc)), varMap.toMap)
def mkFieldReplacer(buf: VarSymbol, baseIdx: VarSymbol, symMap: Map[SimpleSymbol, SimpleSymbol]) =
def getOffset(off: Int)(k: Path => Block): Block =
def getOffset(off: Int, loc: Opt[Loc])(k: Path => Block): Block =
val idxSymbol = new TempSymbol(N, erasedType = S(ErasedType.Int), "idx")
Scoped(Set.single(idxSymbol), Assign(idxSymbol, Call(State.builtinOpsMap("+").asSimpleRef, (baseIdx.asSimpleRef.asArg :: Value.Lit(Tree.IntLit(off)).asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun),
k(DynSelect(buf.asSimpleRef.selSN("buf"), idxSymbol.asSimpleRef, true))))
Scoped(Set.single(idxSymbol), Assign(idxSymbol, Call(State.builtinOpsMap("+").asSimpleRef, (baseIdx.asSimpleRef.asArg :: Value.Lit(Tree.IntLit(off))(N).asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun, N),
k(DynSelect(buf.asSimpleRef.selSN("buf"), idxSymbol.asSimpleRef, true)(loc))))
def assignToOffset(off: Int, r: Result, rst: Block) =
val idxSymbol = new TempSymbol(N, erasedType = S(ErasedType.Int), "idx")
Scoped(Set.single(idxSymbol), Assign(idxSymbol, Call(State.builtinOpsMap("+").asSimpleRef, (baseIdx.asSimpleRef.asArg :: Value.Lit(Tree.IntLit(off)).asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun),
Scoped(Set.single(idxSymbol), Assign(idxSymbol, Call(State.builtinOpsMap("+").asSimpleRef, (baseIdx.asSimpleRef.asArg :: Value.Lit(Tree.IntLit(off))(N).asArg :: Nil) ne_:: Nil)(CallMetadata.defaultMlsFun, N),
AssignDynField(buf.asSimpleRef.selSN("buf"), idxSymbol.asSimpleRef, true, r, applyBlock(rst))))
new BlockTransformer(SymbolSubst.Id):
override def applySimpleSymbol(sym: SimpleSymbol): SimpleSymbol = symMap.getOrElse(sym, sym)
Expand All @@ -66,11 +66,11 @@ class BufferableTransform()(using State, Raise):
case sel: Select =>
sel.symbol.fold(super.applyPath(p)(k)): sym =>
fieldMap.get(sym).orElse(pubFieldMap.get(sym).flatMap(fieldMap.get(_))).fold(super.applyPath(p)(k)): off =>
getOffset(off): res =>
getOffset(off, p.toLoc): res =>
k(res)
case r: Value.Ref =>
fieldMap.get(r.symbol).fold(super.applyPath(p)(k)): off =>
getOffset(off): res =>
getOffset(off, p.toLoc): res =>
k(res)
case _ => super.applyPath(p)(k)
def transformFunDefn(f: FunDefn, isCtor: Bool): FunDefn =
Expand All @@ -79,7 +79,7 @@ class BufferableTransform()(using State, Raise):
val (newParams, symMap) = mkSymbolReplacer(f.params)
val blk = mkFieldReplacer(buf, idx, symMap).applyBlock(f.body)
FunDefn(f.owner, f.sym, TermSymbol(f.dSym.k, f.dSym.owner, f.dSym.id, erasedType = N), PlainParamList(
Param(FldFlags.empty, buf, N, Modulefulness.none) :: Param(FldFlags.empty, idx, N, Modulefulness.none) :: Nil) :: newParams,
Param.simple(buf) :: Param.simple(idx) :: Nil)(N) :: newParams,
if isCtor then Begin(blk, Return(idx.asSimpleRef)) else blk)(configOverride = f.configOverride, annotations = f.annotations)
val fakeCtor = transformFunDefn(FunDefn.withFreshSymbol(
S(companionSym),
Expand All @@ -92,7 +92,7 @@ class BufferableTransform()(using State, Raise):
fakeCtor :: cls.methods.map(transformFunDefn(_, false)),
Nil,
clsSizeSym -> clsSizeTermSym :: Nil,
Define(ValDefn(clsSizeTermSym, clsSizeSym, Value.Lit(Tree.IntLit(fields.size)))(N, Nil), End()),
Define(ValDefn(clsSizeTermSym, clsSizeSym, Value.Lit(Tree.IntLit(fields.size))(N))(N, Nil), End()),
annotations = Nil,
)
k:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,14 +30,14 @@ class ClassParamFlattener(using State) extends BlockTransformer(SymbolSubst.Id):
val flags = paramss.headOption.fold(ParamListFlags.empty)(_.flags)
val (init, last) = paramss.splitAt(paramss.length - 1)
val params = init.flatMap(_.allParams) ::: last.flatMap(_.params)
ParamList(flags, params, last.flatMap(_.restParam).headOption)
ParamList(flags, params, last.flatMap(_.restParam).headOption)(Loc(paramss))

/** Normalize class params so that `paramsOpt = N` and `auxParams` has exactly one element. */
private def flattenClsParams(cls: ClsLikeDefn): ClsLikeDefn =
if cls.paramsOpt.isEmpty && cls.auxParams.sizeIs == 1 then return cls
val paramss = cls.paramsOpt.toList ::: cls.auxParams
val flatAux = paramss match
case Nil => PlainParamList(Nil)
case Nil => PlainParamList(Nil)(N)
case single :: Nil => single
case _ => flattenParamLists(paramss)
cls.copy(paramsOpt = N, auxParams = flatAux :: Nil)(
Expand All @@ -46,8 +46,8 @@ class ClassParamFlattener(using State) extends BlockTransformer(SymbolSubst.Id):
)

private def classPathFor(fun: Path, cls: ClassSymbol): Opt[Path] = fun match
case Value.MemberRef(bms, _) => S(bms.asMemberRef(cls))
case s @ Select(qual, name) => S(Select(qual, name)(S(cls))(s.sanitize))
case Value.MemberRef(bms, _) => S(bms.asMemberRef(cls).withLoc(fun.toLoc))
case s @ Select(qual, name) => S(Select(qual, name)(S(cls), s.toLoc)(s.sanitize))
case _ => N

private def saturatedCurriedClassCall(fun: Path, argss: NELs[Ls[Arg]]): Opt[Path] =
Expand Down Expand Up @@ -77,7 +77,7 @@ class ClassParamFlattener(using State) extends BlockTransformer(SymbolSubst.Id):
else argss2
k:
if flatArgss is argss then c
else Call(r, flatArgss)(c.metadata).withLocOf(c)
else Call(r, flatArgss)(c.metadata, c.toLoc)
case call @ Call(fun, argss) =>
saturatedCurriedClassCall(fun, argss) match
case S(cls) =>
Expand All @@ -86,7 +86,7 @@ class ClassParamFlattener(using State) extends BlockTransformer(SymbolSubst.Id):
val flatArgss =
if argss2.lengthCompare(1) > 0 then argss2.flatten ne_:: Nil
else argss2
k(Instantiate(false, cls2, flatArgss)(InstantiateMetadata(call.metadata.annotations)).withLocOf(call))
k(Instantiate(false, cls2, flatArgss)(InstantiateMetadata(call.metadata.annotations), call.toLoc))
case N =>
super.applyResult(r)(k)
case inst @ Instantiate(mut, cls, argss) =>
Expand All @@ -97,7 +97,7 @@ class ClassParamFlattener(using State) extends BlockTransformer(SymbolSubst.Id):
else argss2
k:
if (cls2 is cls) && (flatArgss is argss) then inst
else Instantiate(mut, cls2, flatArgss)(inst.metadata).withLocOf(inst)
else Instantiate(mut, cls2, flatArgss)(inst.metadata, inst.toLoc)
case _ =>
super.applyResult(r)(k)

Expand Down
Loading
Loading