use crate::quantum::{BaseGate, Circuit, Gate};
use std::f64::consts::PI;
fn hexf(v: f64) -> String {
format!("0x{:016X}", v.to_bits())
}
fn qubit(i: u8) -> String {
format!("%Qubit* inttoptr (i64 {i} to %Qubit*)")
}
fn p_angle(k: u16) -> f64 {
2.0 * PI / (1u64 << k.min(62)) as f64
}
struct Emitter {
body: Vec<String>,
used: Vec<&'static str>,
}
impl Emitter {
fn new() -> Self {
Emitter { body: Vec::new(), used: Vec::new() }
}
fn mark(&mut self, intrinsic: &'static str) {
if !self.used.contains(&intrinsic) {
self.used.push(intrinsic);
}
}
fn one(&mut self, name: &'static str, q: u8) {
self.mark(name);
self.body.push(format!(" call void @{name}({})", qubit(q)));
}
fn two(&mut self, name: &'static str, c: u8, t: u8) {
self.mark(name);
self.body.push(format!(" call void @{name}({}, {})", qubit(c), qubit(t)));
}
fn rz(&mut self, theta: f64, q: u8) {
self.mark("__quantum__qis__rz__body");
self.body.push(format!(
" call void @__quantum__qis__rz__body(double {}, {})",
hexf(theta),
qubit(q)
));
}
fn cy(&mut self, c: u8, t: u8) {
self.one("__quantum__qis__s__adj", t);
self.two("__quantum__qis__cnot__body", c, t);
self.one("__quantum__qis__s__body", t);
}
fn cp(&mut self, theta: f64, c: u8, t: u8) {
self.rz(theta / 2.0, c);
self.two("__quantum__qis__cnot__body", c, t);
self.rz(-theta / 2.0, t);
self.two("__quantum__qis__cnot__body", c, t);
self.rz(theta / 2.0, t);
}
fn ccx(&mut self, a: u8, b: u8, t: u8) {
const H: &str = "__quantum__qis__h__body";
const T: &str = "__quantum__qis__t__body";
const TD: &str = "__quantum__qis__t__adj";
const CN: &str = "__quantum__qis__cnot__body";
self.one(H, t);
self.two(CN, b, t);
self.one(TD, t);
self.two(CN, a, t);
self.one(T, t);
self.two(CN, b, t);
self.one(TD, t);
self.two(CN, a, t);
self.one(T, b);
self.one(T, t);
self.one(H, t);
self.two(CN, a, b);
self.one(T, a);
self.one(TD, b);
self.two(CN, a, b);
}
fn gate(&mut self, g: &Gate) {
let t = g.target;
match (g.base, g.controls.as_slice()) {
(BaseGate::I, _) => {}
(BaseGate::X, []) => self.one("__quantum__qis__x__body", t),
(BaseGate::Y, []) => self.one("__quantum__qis__y__body", t),
(BaseGate::Z, []) => self.one("__quantum__qis__z__body", t),
(BaseGate::H, []) => self.one("__quantum__qis__h__body", t),
(BaseGate::S, []) => self.one("__quantum__qis__s__body", t),
(BaseGate::Sdg, []) => self.one("__quantum__qis__s__adj", t),
(BaseGate::T, []) => self.one("__quantum__qis__t__body", t),
(BaseGate::Tdg, []) => self.one("__quantum__qis__t__adj", t),
(BaseGate::P, []) => self.rz(p_angle(g.param), t),
(BaseGate::X, [c]) => self.two("__quantum__qis__cnot__body", *c, t),
(BaseGate::Z, [c]) => self.two("__quantum__qis__cz__body", *c, t),
(BaseGate::Y, [c]) => self.cy(*c, t),
(BaseGate::P, [c]) => self.cp(p_angle(g.param), *c, t),
(BaseGate::X, [a, b]) => self.ccx(*a, *b, t),
(BaseGate::Z, [a, b]) => {
self.one("__quantum__qis__h__body", t);
self.ccx(*a, *b, t);
self.one("__quantum__qis__h__body", t);
}
(base, ctrls) => self.body.push(format!(
" ; UNMAPPED {base:?} controls={ctrls:?} target={t} — this module is NOT faithful"
)),
}
}
}
fn declaration(name: &str) -> String {
match name {
"__quantum__qis__rz__body" => format!("declare void @{name}(double, %Qubit*)"),
"__quantum__qis__cnot__body" | "__quantum__qis__cz__body" => {
format!("declare void @{name}(%Qubit*, %Qubit*)")
}
"__quantum__qis__mz__body" => format!("declare void @{name}(%Qubit*, %Result*)"),
"__quantum__rt__result_record_output" => format!("declare void @{name}(%Result*, i8*)"),
_ => format!("declare void @{name}(%Qubit*)"),
}
}
pub fn to_qir(c: &Circuit, measure_all: bool) -> String {
let mut e = Emitter::new();
for g in &c.ops {
e.gate(g);
}
let n = c.n_qubits;
if measure_all {
for q in 0..n {
e.mark("__quantum__qis__mz__body");
e.body.push(format!(
" call void @__quantum__qis__mz__body({}, %Result* inttoptr (i64 {q} to %Result*))",
qubit(q)
));
}
for q in 0..n {
e.mark("__quantum__rt__result_record_output");
e.body.push(format!(
" call void @__quantum__rt__result_record_output(%Result* inttoptr (i64 {q} to %Result*), i8* null)"
));
}
}
let mut s = String::new();
s.push_str("; QIR emitted by wai-quantum — deterministic, byte-identical on every target.\n");
s.push_str("source_filename = \"wai_quantum\"\n\n");
s.push_str("%Qubit = type opaque\n%Result = type opaque\n\n");
for i in &e.used {
s.push_str(&declaration(i));
s.push('\n');
}
s.push_str("\ndefine void @main() #0 {\nentry:\n");
for line in &e.body {
s.push_str(line);
s.push('\n');
}
s.push_str(" ret void\n}\n\n");
let results = if measure_all { n } else { 0 };
s.push_str(&format!(
"attributes #0 = {{ \"entry_point\" \"output_labeling_schema\" \
\"qir_profiles\"=\"base_profile\" \"required_num_qubits\"=\"{n}\" \
\"required_num_results\"=\"{results}\" }}\n\n"
));
s.push_str("!llvm.module.flags = !{!0, !1, !2, !3}\n");
s.push_str("!0 = !{i32 1, !\"qir_major_version\", i32 1}\n");
s.push_str("!1 = !{i32 7, !\"qir_minor_version\", i32 0}\n");
s.push_str("!2 = !{i32 1, !\"dynamic_qubit_management\", i1 false}\n");
s.push_str("!3 = !{i32 1, !\"dynamic_result_management\", i1 false}\n");
s
}
fn push(c: &mut Circuit, base: BaseGate, controls: Vec<u8>, target: u8, param: u16) {
c.ops.push(Gate { base, controls, target, param });
}
pub fn rewrite_cy(n: u8, c0: u8, t: u8) -> Circuit {
let mut c = Circuit::new(n);
push(&mut c, BaseGate::Sdg, vec![], t, 0);
c.cx(c0, t);
push(&mut c, BaseGate::S, vec![], t, 0);
c
}
pub fn rewrite_ccx(n: u8, a: u8, b: u8, t: u8) -> Circuit {
let mut c = Circuit::new(n);
c.h(t);
c.cx(b, t);
push(&mut c, BaseGate::Tdg, vec![], t, 0);
c.cx(a, t);
c.t(t);
c.cx(b, t);
push(&mut c, BaseGate::Tdg, vec![], t, 0);
c.cx(a, t);
c.t(b);
c.t(t);
c.h(t);
c.cx(a, b);
c.t(a);
push(&mut c, BaseGate::Tdg, vec![], b, 0);
c.cx(a, b);
c
}
fn angle_to_param(theta: f64) -> Option<u16> {
for k in 0u16..=62 {
let a = p_angle(k);
if (a - theta).abs() <= 1e-9 * a.max(1.0) {
return Some(k);
}
}
None
}
fn arg_indices(line: &str) -> Vec<u8> {
let mut out = Vec::new();
let mut rest = line;
while let Some(i) = rest.find("i64 ") {
rest = &rest[i + 4..];
let end = rest.find(|c: char| !c.is_ascii_digit()).unwrap_or(rest.len());
if let Ok(v) = rest[..end].parse::<u64>() {
out.push(v as u8);
}
rest = &rest[end..];
}
out
}
fn arg_double(line: &str) -> Option<f64> {
let i = line.find("double ")? + 7;
let rest = &line[i..];
let end = rest.find([',', ')']).unwrap_or(rest.len());
let tok = rest[..end].trim();
if let Some(hex) = tok.strip_prefix("0x") {
u64::from_str_radix(hex, 16).ok().map(f64::from_bits)
} else {
tok.parse::<f64>().ok()
}
}
pub fn from_qir(src: &str) -> Result<Circuit, String> {
enum Raw {
G(BaseGate, Vec<u8>, u8),
Rz(f64, u8),
}
let mut raw: Vec<Raw> = Vec::new();
let mut max_q: i32 = -1;
for line in src.lines() {
let line = line.trim();
if !line.starts_with("call void @__quantum__") {
continue;
}
let name = line
.split('@')
.nth(1)
.and_then(|r| r.split('(').next())
.ok_or_else(|| format!("malformed call: {line}"))?;
let qs = arg_indices(line);
for q in &qs {
max_q = max_q.max(*q as i32);
}
let need = |n: usize| -> Result<(), String> {
if qs.len() < n {
Err(format!("{name}: expected {n} qubit operands, got {}", qs.len()))
} else {
Ok(())
}
};
match name {
"__quantum__qis__x__body" => { need(1)?; raw.push(Raw::G(BaseGate::X, vec![], qs[0])); }
"__quantum__qis__y__body" => { need(1)?; raw.push(Raw::G(BaseGate::Y, vec![], qs[0])); }
"__quantum__qis__z__body" => { need(1)?; raw.push(Raw::G(BaseGate::Z, vec![], qs[0])); }
"__quantum__qis__h__body" => { need(1)?; raw.push(Raw::G(BaseGate::H, vec![], qs[0])); }
"__quantum__qis__s__body" => { need(1)?; raw.push(Raw::G(BaseGate::S, vec![], qs[0])); }
"__quantum__qis__s__adj" => { need(1)?; raw.push(Raw::G(BaseGate::Sdg, vec![], qs[0])); }
"__quantum__qis__t__body" => { need(1)?; raw.push(Raw::G(BaseGate::T, vec![], qs[0])); }
"__quantum__qis__t__adj" => { need(1)?; raw.push(Raw::G(BaseGate::Tdg, vec![], qs[0])); }
"__quantum__qis__cnot__body" | "__quantum__qis__cx__body" => {
need(2)?; raw.push(Raw::G(BaseGate::X, vec![qs[0]], qs[1]));
}
"__quantum__qis__cz__body" => { need(2)?; raw.push(Raw::G(BaseGate::Z, vec![qs[0]], qs[1])); }
"__quantum__qis__rz__body" => {
need(1)?;
let theta = arg_double(line).ok_or_else(|| format!("rz: no double operand: {line}"))?;
raw.push(Raw::Rz(theta, qs[0]));
}
"__quantum__qis__mz__body" | "__quantum__qis__m__body"
| "__quantum__rt__result_record_output" | "__quantum__rt__tuple_record_output"
| "__quantum__rt__array_record_output" => {}
other => return Err(format!("unsupported intrinsic: {other}")),
}
}
let close = |x: f64, y: f64| (x - y).abs() <= 1e-9 * x.abs().max(1.0);
let mut gates: Vec<(BaseGate, Vec<u8>, u8, u16)> = Vec::new();
let mut i = 0usize;
while i < raw.len() {
if i + 4 < raw.len()
&& let (
Raw::Rz(a0, c0),
Raw::G(BaseGate::X, cs1, t1),
Raw::Rz(a2, t2),
Raw::G(BaseGate::X, cs3, t3),
Raw::Rz(a4, t4),
) = (&raw[i], &raw[i + 1], &raw[i + 2], &raw[i + 3], &raw[i + 4])
&& cs1.len() == 1
&& cs3.len() == 1
&& cs1[0] == *c0
&& cs3[0] == *c0
&& t1 == t2
&& t1 == t3
&& t1 == t4
&& close(*a2, -*a0)
&& close(*a4, *a0)
&& let Some(k) = angle_to_param(2.0 * *a0)
{
gates.push((BaseGate::P, vec![*c0], *t1, k));
i += 5;
continue;
}
match &raw[i] {
Raw::G(b, cs, t) => gates.push((*b, cs.clone(), *t, 0)),
Raw::Rz(theta, q) => {
let k = angle_to_param(*theta).ok_or_else(|| {
format!("rz({theta}) is not a representable 2π/2^k phase — this gate set stores dyadic phases only")
})?;
gates.push((BaseGate::P, vec![], *q, k));
}
}
i += 1;
}
let declared = src
.find("\"required_num_qubits\"=\"")
.and_then(|i| src[i + 23..].split('"').next().and_then(|v| v.parse::<u8>().ok()));
let n = declared.unwrap_or_else(|| (max_q + 1).max(1) as u8);
if let Some(d) = declared
&& max_q + 1 > d as i32
{
return Err(format!("module uses qubit {max_q} but declares required_num_qubits={d}"));
}
let mut c = Circuit::new(n);
for (base, controls, target, param) in gates {
c.ops.push(Gate { base, controls, target, param });
}
c.validate().map_err(|e| format!("parsed an invalid circuit: {e:?}"))?;
Ok(c)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantum::ONE;
use crate::quantum_toolchain as qt;
fn agrees_on_all_inputs(n: u8, build: &dyn Fn(&mut Circuit), rewrite: &Circuit) -> (bool, i64) {
let tol = ONE / 100_000; let mut worst = ONE;
for k in 0..(1usize << n) {
let mut a = Circuit::new(n);
let mut b = Circuit::new(n);
for q in 0..n {
if k >> q & 1 == 1 {
a.x(q);
b.x(q);
}
}
build(&mut a);
for g in &rewrite.ops {
b.ops.push(g.clone());
}
let f = a.simulate().unwrap().fidelity_fx(&b.simulate().unwrap());
worst = worst.min(f);
if (ONE - f).abs() > tol {
return (false, f);
}
}
(true, worst)
}
#[test]
fn cy_rewrite_is_exact() {
let (ok, f) = agrees_on_all_inputs(
2,
&|c: &mut Circuit| {
c.ops.push(Gate { base: BaseGate::Y, controls: vec![0], target: 1, param: 0 });
},
&rewrite_cy(2, 0, 1),
);
assert!(ok, "CY rewrite fidelity {f} of {ONE}");
}
#[test]
fn ccx_rewrite_is_exact() {
let (ok, f) = agrees_on_all_inputs(3, &|c: &mut Circuit| { c.ccx(0, 1, 2); }, &rewrite_ccx(3, 0, 1, 2));
assert!(ok, "CCX rewrite fidelity {f} of {ONE}");
}
#[test]
fn ccx_rewrite_holds_with_controls_swapped() {
let (ok, f) = agrees_on_all_inputs(3, &|c: &mut Circuit| { c.ccx(1, 0, 2); }, &rewrite_ccx(3, 1, 0, 2));
assert!(ok, "CCX(swapped) rewrite fidelity {f} of {ONE}");
}
type C = (f64, f64);
fn cmul(a: C, b: C) -> C { (a.0 * b.0 - a.1 * b.1, a.0 * b.1 + a.1 * b.0) }
fn expi(t: f64) -> C { (t.cos(), t.sin()) }
fn rz_on(v: &mut [C; 4], phi: f64, q: usize) {
for (i, e) in v.iter_mut().enumerate() {
let bit = if q == 0 { i >> 1 & 1 } else { i & 1 };
let sign = if bit == 1 { 1.0 } else { -1.0 };
*e = cmul(*e, expi(sign * phi / 2.0));
}
}
fn cnot(v: &mut [C; 4]) {
v.swap(0b10, 0b11);
}
#[test]
fn cp_mapping_is_exact_up_to_global_phase() {
for &k in &[1u16, 2, 3, 5, 8] {
let theta = p_angle(k);
let mut global: Option<C> = None;
for basis in 0..4usize {
let mut v: [C; 4] = [(0.0, 0.0); 4];
v[basis] = (1.0, 0.0);
rz_on(&mut v, theta / 2.0, 0);
cnot(&mut v);
rz_on(&mut v, -theta / 2.0, 1);
cnot(&mut v);
rz_on(&mut v, theta / 2.0, 1);
let want: C = if basis == 3 { expi(theta) } else { (1.0, 0.0) };
let got = v[basis];
let ratio = cmul(got, (want.0, -want.1)); match global {
None => global = Some(ratio),
Some(g) => {
assert!((g.0 - ratio.0).abs() < 1e-12 && (g.1 - ratio.1).abs() < 1e-12,
"k={k} basis={basis}: phase {ratio:?} != {g:?} — not a GLOBAL phase");
}
}
for (j, e) in v.iter().enumerate() {
if j != basis {
assert!(e.0.abs() < 1e-12 && e.1.abs() < 1e-12, "k={k} leaked into {j}");
}
}
}
let g = global.unwrap();
assert!((g.0 * g.0 + g.1 * g.1 - 1.0).abs() < 1e-12, "global phase must be unit");
}
}
#[test]
fn emits_a_wellformed_module() {
let c = qt::ghz(3);
let ir = to_qir(&c, true);
assert!(ir.contains("%Qubit = type opaque"));
assert!(ir.contains("declare void @__quantum__qis__h__body(%Qubit*)"));
assert!(ir.contains("declare void @__quantum__qis__cnot__body(%Qubit*, %Qubit*)"));
assert!(ir.contains("define void @main() #0 {"));
assert!(ir.contains("\"required_num_qubits\"=\"3\""));
assert!(ir.contains("\"required_num_results\"=\"3\""));
assert!(ir.contains("qir_major_version"));
assert_eq!(ir.matches("__quantum__qis__cnot__body(%Qubit* inttoptr").count(), 2);
assert!(!ir.contains("UNMAPPED"));
}
#[test]
fn angles_are_exact_hex_floats() {
let mut c = Circuit::new(1);
c.p(3, 0); let ir = to_qir(&c, false);
let want = hexf(std::f64::consts::FRAC_PI_4);
assert!(ir.contains(&format!("double {want}")), "{ir}");
let bits = u64::from_str_radix(&want[2..], 16).unwrap();
assert_eq!(f64::from_bits(bits), std::f64::consts::FRAC_PI_4);
}
#[test]
fn every_toolchain_algorithm_maps_completely() {
for (id, _, _) in qt::catalog() {
if let Some(c) = qt::build_algorithm(id, 3, 5) {
let ir = to_qir(&c, true);
assert!(!ir.contains("UNMAPPED"), "{id} produced an unfaithful module:\n{ir}");
}
}
}
#[test]
fn no_symbol_is_declared_twice() {
for (id, _, _) in qt::catalog() {
if let Some(c) = qt::build_algorithm(id, 3, 5) {
for measure in [true, false] {
let ir = to_qir(&c, measure);
let mut names: Vec<&str> = ir
.lines()
.filter(|l| l.starts_with("declare "))
.map(|l| l.split('@').nth(1).unwrap().split('(').next().unwrap())
.collect();
let n = names.len();
names.sort_unstable();
names.dedup();
assert_eq!(n, names.len(), "{id}: duplicate declare in\n{ir}");
}
}
}
}
#[test]
fn round_trips_semantically() {
let mut checked = 0;
for (id, _, _) in qt::catalog() {
let Some(c) = qt::build_algorithm(id, 3, 5) else { continue };
let ir = to_qir(&c, true);
match from_qir(&ir) {
Ok(back) => {
assert_eq!(back.n_qubits, c.n_qubits, "{id}: width changed");
let f = c.simulate().unwrap().fidelity_fx(&back.simulate().unwrap());
assert!(
(ONE - f).abs() <= ONE / 100_000,
"{id}: round-trip fidelity {f} of {ONE}"
);
checked += 1;
}
Err(e) => {
assert!(e.contains("not a representable"), "{id}: unexpected error: {e}");
}
}
}
assert!(checked >= 5, "expected most algorithms to round-trip, got {checked}");
}
#[test]
fn qft_round_trips_to_an_identical_statevector() {
for n in 2..=4u8 {
let c = qt::build_algorithm("qft", n, 0).unwrap();
let back = from_qir(&to_qir(&c, true)).expect("QFT must ingest");
assert_eq!(
c.simulate().unwrap().statevector_hash(),
back.simulate().unwrap().statevector_hash(),
"qft n={n} did not return identical"
);
assert!(
c.ops.iter().any(|g| matches!(g.base, BaseGate::P) && !g.controls.is_empty()),
"qft n={n} should contain a controlled phase for this to be meaningful"
);
}
}
#[test]
fn refuses_unrepresentable_angles_instead_of_approximating() {
let ir = format!(
"define void @main() #0 {{\nentry:\n call void @__quantum__qis__rz__body(double {}, %Qubit* inttoptr (i64 0 to %Qubit*))\n ret void\n}}\nattributes #0 = {{ \"required_num_qubits\"=\"1\" }}\n",
hexf(0.1)
);
let e = from_qir(&ir).unwrap_err();
assert!(e.contains("not a representable"), "{e}");
}
#[test]
fn rejects_unknown_intrinsics_and_overwide_modules() {
let bad = "define void @main() #0 {\nentry:\n call void @__quantum__qis__toffoli__body(%Qubit* inttoptr (i64 0 to %Qubit*))\n ret void\n}\n";
assert!(from_qir(bad).unwrap_err().contains("unsupported intrinsic"));
let over = "define void @main() #0 {\nentry:\n call void @__quantum__qis__h__body(%Qubit* inttoptr (i64 5 to %Qubit*))\n ret void\n}\nattributes #0 = { \"required_num_qubits\"=\"2\" }\n";
assert!(from_qir(over).unwrap_err().contains("required_num_qubits"));
}
#[test]
fn ingests_a_hand_written_clifford_t_module() {
let ir = "\
define void @main() {
entry:
call void @__quantum__qis__h__body(%Qubit* inttoptr (i64 0 to %Qubit*))
call void @__quantum__qis__cnot__body(%Qubit* inttoptr (i64 0 to %Qubit*), %Qubit* inttoptr (i64 1 to %Qubit*))
call void @__quantum__qis__t__body(%Qubit* inttoptr (i64 1 to %Qubit*))
call void @__quantum__qis__mz__body(%Qubit* inttoptr (i64 0 to %Qubit*), %Result* inttoptr (i64 0 to %Result*))
ret void
}
";
let c = from_qir(ir).expect("should ingest");
assert_eq!(c.n_qubits, 2, "width inferred from operands");
assert_eq!(c.ops.len(), 3, "measurement is not a gate");
let mut want = Circuit::new(2);
want.h(0);
want.cx(0, 1);
want.t(1);
assert_eq!(
want.simulate().unwrap().statevector_hash(),
c.simulate().unwrap().statevector_hash()
);
}
#[test]
fn deterministic() {
let c = qt::grover(3, 5);
assert_eq!(to_qir(&c, true), to_qir(&c, true));
}
}