use super::super::super::field::Field;
use std::collections::HashMap;
#[cfg(test)]
mod tests;
#[derive(Clone, Copy, Debug)]
pub enum ConnectionType<T>
where
T: Copy,
{
Left(T, SubCircuitId),
Right(T, SubCircuitId),
Output(SubCircuitId),
}
#[derive(Clone)]
pub struct SubCircuitConnections<T> {
left_inputs: Vec<(T, WireId)>,
right_inputs: Vec<(T, WireId)>,
output: WireId,
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct WireId(usize);
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct SubCircuitId(usize);
impl SubCircuitId {
pub fn inner_id(&self) -> usize {
self.0
}
}
pub struct Circuit<T>
where
T: Copy,
{
next_wire_id: WireId,
next_sub_circuit_id: SubCircuitId,
wire_assignments: HashMap<WireId, Vec<ConnectionType<T>>>,
sub_circuit_wires: HashMap<SubCircuitId, SubCircuitConnections<T>>,
wire_values: HashMap<WireId, Option<T>>,
}
impl<T> Circuit<T>
where
T: Copy + Field,
{
pub fn new() -> Self {
Circuit {
next_wire_id: WireId(1),
next_sub_circuit_id: SubCircuitId(0),
wire_assignments: HashMap::new(),
sub_circuit_wires: HashMap::new(),
wire_values: HashMap::new(),
}
}
pub fn unity_wire(&self) -> WireId {
WireId(0)
}
pub fn new_wire(&mut self) -> WireId {
let next_wire_id = self.next_wire_id;
self.next_wire_id.0 += 1;
self.wire_values.insert(next_wire_id, None);
next_wire_id
}
pub fn num_wires(&self) -> usize {
self.next_wire_id.0
}
pub fn value(&self, wire: WireId) -> Option<T> {
*self
.wire_values
.get(&wire)
.expect("wire is not defined in this circuit")
}
pub fn set_value(&mut self, wire: WireId, value: T) {
self.wire_values.insert(wire, Some(value));
}
pub fn wire_assignments(&self) -> &HashMap<WireId, Vec<ConnectionType<T>>> {
&self.wire_assignments
}
pub fn assignments(&self, wire: &WireId) -> &Vec<ConnectionType<T>> {
self.wire_assignments
.get(wire)
.expect("wire id is not defined in this circuit")
}
fn insert_connection(&mut self, wire: WireId, connection: ConnectionType<T>) {
if self.wire_assignments.get(&wire).is_none() {
self.wire_assignments.insert(wire, vec![connection]);
} else {
self.wire_assignments
.get_mut(&wire)
.map(|v| v.push(connection));
}
}
pub fn sub_circuits(&self) -> impl Iterator<Item = SubCircuitId> {
(0..self.next_sub_circuit_id.0).map(|id| SubCircuitId(id))
}
pub fn new_sub_circuit(
&mut self,
left_inputs: Vec<(T, WireId)>,
right_inputs: Vec<(T, WireId)>,
) -> WireId {
use self::ConnectionType::{Left, Output, Right};
let sub_circuit_id = self.next_sub_circuit_id;
self.next_sub_circuit_id.0 += 1;
let output_wire = self.new_wire();
for (weight, wire) in left_inputs.clone().into_iter() {
let connection = Left(weight, sub_circuit_id);
self.insert_connection(wire, connection);
}
for (weight, wire) in right_inputs.clone().into_iter() {
let connection = Right(weight, sub_circuit_id);
self.insert_connection(wire, connection);
}
let connection = Output(sub_circuit_id);
self.insert_connection(output_wire, connection);
self.sub_circuit_wires.insert(
sub_circuit_id,
SubCircuitConnections {
left_inputs,
right_inputs,
output: output_wire,
},
);
output_wire
}
fn evaluate_sub_circuit(&mut self, sub_circuit: SubCircuitId) -> T {
let SubCircuitConnections {
left_inputs,
right_inputs,
..
} = self
.sub_circuit_wires
.get(&sub_circuit)
.expect("a sub circuit referenced by a wire should exist")
.clone();
let lhs = left_inputs
.into_iter()
.fold(T::zero(), |acc, (weight, wire)| {
acc + weight * self.evaluate(wire)
});
let rhs = right_inputs
.into_iter()
.fold(T::zero(), |acc, (weight, wire)| {
acc + weight * self.evaluate(wire)
});
lhs * rhs
}
pub fn evaluate(&mut self, wire: WireId) -> T {
use self::ConnectionType::Output;
if wire == self.unity_wire() {
return T::one();
}
self.wire_values
.get(&wire)
.expect("cannot evaluate unknown wire")
.unwrap_or_else(|| {
let output_sub_circuit = self
.wire_assignments
.get(&wire)
.expect("a wire must be attached to something")
.into_iter()
.filter_map(|c| if let &Output(sc) = c { Some(sc) } else { None })
.nth(0)
.expect("a wire with an unknown value must be the output of a sub circuit");
let value = self.evaluate_sub_circuit(output_sub_circuit);
self.wire_values.insert(wire, Some(value));
value
})
}
pub fn reset(&mut self) {
for value in self.wire_values.values_mut() {
*value = None;
}
}
pub fn new_bit_checker(&mut self, input: WireId) -> WireId {
let lhs_inputs = vec![(T::one(), input)];
let rhs_inputs = vec![(T::one(), input), (-T::one(), self.unity_wire())];
self.new_sub_circuit(lhs_inputs, rhs_inputs)
}
pub fn new_not(&mut self, input: WireId) -> WireId {
let lhs_inputs = vec![(T::one(), self.unity_wire())];
let rhs_inputs = vec![(T::one(), self.unity_wire()), (-T::one(), input)];
self.new_sub_circuit(lhs_inputs, rhs_inputs)
}
pub fn new_and(&mut self, lhs: WireId, rhs: WireId) -> WireId {
let lhs_inputs = vec![(T::one(), lhs)];
let rhs_inputs = vec![(T::one(), rhs)];
self.new_sub_circuit(lhs_inputs, rhs_inputs)
}
pub fn new_or(&mut self, lhs: WireId, rhs: WireId) -> WireId {
let lhs_and_rhs = self.new_and(lhs, rhs);
let one = T::one();
let lhs_inputs = vec![(-one, lhs_and_rhs), (one, lhs), (one, rhs)];
let rhs_inputs = vec![(one, self.unity_wire())];
self.new_sub_circuit(lhs_inputs, rhs_inputs)
}
pub fn new_xor(&mut self, lhs: WireId, rhs: WireId) -> WireId {
let one = T::one();
let lhs_inputs = vec![(one, lhs), (-one, rhs)];
let rhs_inputs = vec![(one, lhs), (-one, rhs)];
self.new_sub_circuit(lhs_inputs, rhs_inputs)
}
pub fn fan_in<F>(&mut self, inputs: &[WireId], mut gate: F) -> WireId
where
F: FnMut(&mut Self, WireId, WireId) -> WireId,
{
if inputs.len() < 2 {
panic!("cannot fan in with fewer than two inputs");
}
inputs
.iter()
.skip(1)
.fold(inputs[0], |acc, wire| gate(self, acc, *wire))
}
pub fn bitwise_op<F>(&mut self, left: &[WireId], right: &[WireId], mut gate: F) -> Vec<WireId>
where
F: FnMut(&mut Self, WireId, WireId) -> WireId,
{
assert!(left.len() == right.len());
left.iter()
.zip(right.iter())
.map(|(&l, &r)| gate(self, l, r))
.collect()
}
}