diff --git a/tools/src/dafny-emit.ts b/tools/src/dafny-emit.ts index 9d600ed..e471dc8 100644 --- a/tools/src/dafny-emit.ts +++ b/tools/src/dafny-emit.ts @@ -306,15 +306,25 @@ function emitExpr(e: Expr): string { needPreamble("SeqJoin"); return `SeqJoin(${obj}, ${args[0]})`; } - if (e.method === "some" && e.args[0].kind === "lambda" && - e.args[0].body.length === 1 && e.args[0].body[0].kind === "return") { + // `.some(pred)`: inline a single-return lambda's body, else apply the + // predicate (e.g. a function reference). + if (e.method === "some") { const lam = e.args[0]; - const ret = lam.body[0]; - if (ret.kind !== "return") throw new Error("unreachable"); - const { binder: p, body: v } = comprehensionBinder(lam, ret.value, e.obj); - const body = emitExpr(v); + let p: string, body: string; + if (lam.kind === "lambda" && lam.body.length === 1 && lam.body[0].kind === "return") { + const cb = comprehensionBinder(lam, lam.body[0].value, e.obj); + p = cb.binder; + body = emitExpr(cb.body); + } else { + p = escapeName(freshBinder("x", e.obj, e.args[0])); + body = `${args[0]}(${p})`; + } return `(exists ${p} :: ${p} in ${obj} && ${body})`; } + // `.reduce(f, init)` → Std's FoldLeft(f, init, xs) (same arg order). + if (e.method === "reduce" && args.length === 2) { + return `Std.Collections.Seq.FoldLeft(${args[0]}, ${args[1]}, ${obj})`; + } } // String methods if (ty === "string") { diff --git a/tools/src/lean-emit.ts b/tools/src/lean-emit.ts index cb62c30..dbdc4b9 100644 --- a/tools/src/lean-emit.ts +++ b/tools/src/lean-emit.ts @@ -237,6 +237,7 @@ function emitMethodCall(tyKind: string, method: string, monadic: boolean, obj: s if (method === "filter") return `${obj}.${monadic ? "filterM" : "filter"} ${args[0]}`; if (method === "every") return `${obj}.${monadic ? "allM" : "all"} ${args[0]}`; if (method === "some") return `${obj}.${monadic ? "anyM" : "any"} ${args[0]}`; + if (method === "reduce" && args.length === 2) return `(${obj}.foldl ${args[0]} ${args[1]})`; if (method === "includes") return args.length > 1 ? `(${obj}.extract ${args[1]} ${obj}.size).contains ${args[0]}` : `${obj}.contains ${args[0]}`; if (method === "find") return `${obj}.find? ${args[0]}`; if (method === "join") return `(String.intercalate ${args[0]} ${obj}.toList)`; diff --git a/tools/src/resolve.ts b/tools/src/resolve.ts index 5209a93..08ad755 100644 --- a/tools/src/resolve.ts +++ b/tools/src/resolve.ts @@ -495,6 +495,18 @@ function inferLambdaParamTypes(fn: TExpr, rawArgs: RawExpr[], ctx?: Ctx): RawExp return [{ ...lam, params: updatedParams }, ...rawArgs.slice(1)]; } } + // reduce's callback is (acc, elem): acc from the init arg's type, elem from the array. + if (fn.kind === "field" && fn.obj.ty.kind === "array" && fn.field === "reduce" && ctx && + rawArgs.length >= 2 && rawArgs[0].kind === "lambda" && rawArgs[0].params.length >= 2) { + const accTs = tyToTsStr(resolveExpr(rawArgs[1], ctx).ty); + const elemTs = tyToTsStr(fn.obj.ty.elem); + if (accTs && elemTs) { + const lam = rawArgs[0]; + const updatedParams = lam.params.map((p, i) => + p.tsType || i > 1 ? p : { ...p, tsType: i === 0 ? accTs : elemTs }); + return [{ ...lam, params: updatedParams }, ...rawArgs.slice(1)]; + } + } if (fn.kind === "field" && fn.obj.ty.kind === "array" && ["map", "filter", "every", "some", "find", "findLast", "findIndex", "findLastIndex"].includes(fn.field) && rawArgs.length >= 1 && rawArgs[0].kind === "lambda" && @@ -597,6 +609,7 @@ function inferMethodReturnTy(fn: TExpr, args: TExpr[], ctx: Ctx): Ty { if (fn.field === "sort") return objTy; if (fn.field === "filter") return objTy; if (fn.field === "every" || fn.field === "some") return { kind: "bool" }; + if (fn.field === "reduce" && args.length === 2) return args[1].ty; if (fn.field === "find" || fn.field === "findLast") return { kind: "optional", inner: objTy.elem }; if (fn.field === "findIndex" || fn.field === "findLastIndex") return { kind: "int" }; if (fn.field === "flat" && objTy.elem.kind === "array") return { kind: "array", elem: objTy.elem.elem };