14#include "r1cs/Dialect/IR/Ops.h"
15#include "r1cs/Transforms/TransformationPassPipelines.h"
22#include <mlir/Pass/PassManager.h>
23#include <mlir/Support/FileUtilities.h>
25#include <llvm/ADT/SmallVector.h>
26#include <llvm/ADT/StringExtras.h>
27#include <llvm/Support/ToolOutputFile.h>
38constexpr char WTNS_MAGIC[] = {
'w',
't',
'n',
's'};
39constexpr uint32_t WTNS_VERSION = 2;
40constexpr uint32_t WTNS_SECTION_COUNT = 2;
41constexpr uint32_t WTNS_HEADER_SECTION = 1;
42constexpr uint32_t WTNS_VALUES_SECTION = 2;
43constexpr uint32_t WTNS_FIELD_LIMB_BITS = 64;
44constexpr uint32_t WTNS_FIELD_LIMB_BYTES = WTNS_FIELD_LIMB_BITS / CHAR_BIT;
46void appendSection(BinaryBuffer &
file, uint32_t type,
const BinaryBuffer §ion) {
48 file.writeU64(section.size());
49 file.writeBytes(section.bytes());
52llvm::Expected<llvm::DynamicAPInt>
53readFelt(
const llvm::json::Object &
object, StringRef key,
const Field &field) {
54 const llvm::json::Value *json =
object.get(key);
56 return makeError(llvm::Twine(
"full witness is missing '") + key +
"'");
58 std::optional<StringRef> text = json->getAsString();
59 if (!text || text->empty() || !llvm::all_of(*text, llvm::isDigit)) {
60 return makeError(llvm::Twine(
"full witness value '") + key +
"' is not a non-negative integer");
63 if (value >= field.prime()) {
64 return makeError(llvm::Twine(
"full witness value '") + key +
"' is outside the field");
69llvm::Expected<OwningOpRef<ModuleOp>> lowerModuleToR1CS(ModuleOp moduleOp) {
70 OwningOpRef<ModuleOp> lowered = cast<ModuleOp>(moduleOp->clone());
71 PassManager pm(lowered->getContext());
72 r1cs::buildFullR1CSLoweringPipeline(pm);
73 if (failed(pm.run(*lowered))) {
74 return makeError(
"failed to lower a module clone while validating .wtns wire ordering");
79llvm::Expected<size_t> getR1CSWireCount(ModuleOp r1csModule, StringRef circuitName) {
80 auto circuit = r1csModule.lookupSymbol<r1cs::CircuitDefOp>(circuitName);
82 return makeError(
"R1CS lowering did not produce a circuit for the llzk.main struct");
85 Block &entry = circuit.getBody().front();
86 size_t wireCount = 1 + entry.getNumArguments();
87 wireCount += llvm::range_size(entry.getOps<r1cs::SignalDefOp>());
91llvm::Expected<std::string> getMainStructName(ModuleOp moduleOp) {
92 SymbolTableCollection tables;
94 if (failed(mainDef) || !mainDef.value()) {
95 return makeError(
"module is missing a concrete llzk.main struct");
97 return mainDef->get().getSymName().str();
100llvm::Expected<SmallVector<llvm::DynamicAPInt>>
101collectWitnessValues(ModuleOp moduleOp,
const llvm::json::Value &fullWitness,
const Field &field) {
102 const auto *root = fullWitness.getAsObject();
103 const auto *inputs = root ? root->getObject(
"inputs") :
nullptr;
104 const auto *signals = root ? root->getObject(
"signals") :
nullptr;
105 if (!inputs || !signals) {
106 return makeError(
".wtns output requires a full-witness llzk-witgen result");
109 SymbolTableCollection tables;
111 if (failed(mainDef) || !mainDef.value()) {
112 return makeError(
"module is missing a concrete llzk.main struct");
114 auto compute = mainDef->get().getComputeFuncOp();
115 auto constrain = mainDef->get().getConstrainFuncOp();
117 return makeError(
"main struct is missing @compute");
119 if (!constrain || constrain.getNumArguments() != compute.getNumArguments() + 1) {
120 return makeError(
"main @constrain inputs do not match @compute inputs");
123 SmallVector<llvm::DynamicAPInt> witness;
126 witness.push_back(field.one());
127 auto appendMemberClass = [&](
bool isPublic) -> llvm::Error {
128 for (component::MemberDefOp member : mainDef->get().getMemberDefs()) {
129 if (member.hasPublicAttr() != isPublic) {
132 if (!isa<felt::FeltType>(member.getType())) {
133 return makeError(
".wtns output currently requires scalar felt main members");
135 auto value = readFelt(*signals, member.getSymName(), field);
137 return value.takeError();
139 witness.push_back(*value);
141 return llvm::Error::success();
144 auto appendInputClass = [&](
bool isPublic) -> llvm::Error {
148 if (constrain.hasArgPublicAttr(binding.index + 1) != isPublic) {
151 if (!isa<felt::FeltType>(binding.type)) {
152 return makeError(
".wtns output currently requires scalar felt main inputs");
154 auto value = readFelt(*inputs, binding.name, field);
156 return value.takeError();
158 witness.push_back(*value);
160 return llvm::Error::success();
163 if (
auto error = appendMemberClass(
true)) {
166 if (
auto error = appendInputClass(
true)) {
169 if (
auto error = appendInputClass(
false)) {
172 if (
auto error = appendMemberClass(
false)) {
179llvm::Expected<BinaryBuffer>
180serializeWtns(ArrayRef<llvm::DynamicAPInt> witness,
const Field &field) {
181 if (witness.size() > std::numeric_limits<uint32_t>::max()) {
182 return makeError(
"witness length does not fit in the .wtns header");
184 uint32_t fieldSize = ((field.bitWidth() + WTNS_FIELD_LIMB_BITS - 1) / WTNS_FIELD_LIMB_BITS) *
185 WTNS_FIELD_LIMB_BYTES;
188 header.writeU32(fieldSize);
189 header.writeFieldElement(fieldSize, field.prime());
190 header.writeU32(
static_cast<uint32_t
>(witness.size()));
193 for (
const llvm::DynamicAPInt &value : witness) {
194 values.writeFieldElement(fieldSize, value);
198 file.writeBytes(WTNS_MAGIC);
199 file.writeU32(WTNS_VERSION);
200 file.writeU32(WTNS_SECTION_COUNT);
201 appendSection(
file, WTNS_HEADER_SECTION, header);
202 appendSection(
file, WTNS_VALUES_SECTION, values);
206llvm::Error writeBinaryFile(StringRef outputFilename,
const BinaryBuffer &
file) {
207 std::unique_ptr<llvm::ToolOutputFile> output = openOutputFile(outputFilename);
209 return makeError(llvm::Twine(
"failed to open .wtns output: ") + outputFilename);
211 output->os().write(
file.bytes().data(),
file.bytes().size());
212 output->os().flush();
213 if (output->os().has_error()) {
214 return makeError(llvm::Twine(
"failed to write .wtns output: ") + outputFilename);
217 return llvm::Error::success();
223 ModuleOp moduleOp,
const llvm::json::Value &fullWitness,
const Field &field,
224 StringRef outputFilename
226 auto mainName = getMainStructName(moduleOp);
228 return mainName.takeError();
231 auto witness = collectWitnessValues(moduleOp, fullWitness, field);
233 return witness.takeError();
236 auto loweredModule = lowerModuleToR1CS(moduleOp);
237 if (!loweredModule) {
238 return loweredModule.takeError();
240 auto r1csWireCount = getR1CSWireCount(**loweredModule, *mainName);
241 if (!r1csWireCount) {
242 return r1csWireCount.takeError();
246 if (*r1csWireCount != witness->size()) {
248 llvm::Twine(
"cannot emit .wtns: R1CS lowering produces ") + llvm::Twine(*r1csWireCount) +
249 " wires but llzk-witgen collected " + llvm::Twine(witness->size()) +
250 "; synthesized R1CS auxiliary wires are not yet supported"
254 auto file = serializeWtns(*witness, field);
256 return file.takeError();
259 return writeBinaryFile(outputFilename, *
file);
This file implements helper methods for constructing DynamicAPInts.
then any Derivative Works that You distribute must include a readable copy of the attribution notices contained within such NOTICE file
Information about the prime finite field used for the interval analysis.
llvm::SmallVector< InputBinding > collectInputBindings(function::FuncDefOp computeFunc)
Collect stable JSON bindings for the main compute inputs.
llvm::Error writeWtns(ModuleOp moduleOp, const llvm::json::Value &fullWitness, const Field &field, StringRef outputFilename)
llvm::Error makeError(const llvm::Twine &msg)
Build a string-backed error for user-facing witgen failures.
DynamicAPInt toDynamicAPInt(StringRef str)
FailureOr< SymbolLookupResult< StructDefOp > > getMainInstanceDef(SymbolTableCollection &symbolTable, Operation *lookupFrom)