Skip to main content

year2020/
day14.rs

1use utils::bit::BitIterator;
2use utils::hash::FastSet;
3use utils::prelude::*;
4
5/// Applying bitmasks to values and memory addresses.
6#[derive(Clone, Debug)]
7pub struct Day14 {
8    writes: Vec<Write>,
9}
10
11#[derive(Copy, Clone, Debug)]
12struct Write {
13    address: u64,
14    value: u64,
15    ones: u64,
16    floating: u64,
17}
18
19#[derive(Copy, Clone, Debug)]
20struct AddressSet {
21    ones: u64,
22    floating: u64,
23}
24
25const VALUE_MASK: u64 = (1 << 36) - 1;
26const INDEX_BITS: usize = 8;
27const INDEX_MASK: u64 = (1 << INDEX_BITS) - 1;
28
29impl Day14 {
30    pub fn new(input: &str, _: InputType) -> Result<Self, InputError> {
31        enum Instruction {
32            Mask((u64, u64)),
33            Mem(u64, u64),
34        }
35
36        if input.is_empty() {
37            return Err(InputError::new(input, 0, "expected instruction"));
38        }
39
40        let mask = parser::byte_map!(b'0' => (0, 0), b'1' => (1, 0), b'X' => (0, 1))
41            .repeat_n::<36, _>(parser::noop())
42            .map(|bits| {
43                bits.iter().fold((0, 0), |(ones, floating), &(one, x)| {
44                    ((ones << 1) | one, (floating << 1) | x)
45                })
46            });
47        let number = parser::number_range(0..=VALUE_MASK);
48        let instruction = parser::parse_tree!(
49            ("mask = ", mask @ mask) => Instruction::Mask(mask),
50            ("mem[", address @ number, "] = ", value @ number) => Instruction::Mem(address, value),
51        );
52
53        let mut writes = Vec::new();
54        let mut current_mask = None;
55        for item in instruction.with_eol().parse_iterator(input) {
56            match item? {
57                Instruction::Mask(mask) => current_mask = Some(mask),
58                Instruction::Mem(address, value) => {
59                    let Some((ones, floating)) = current_mask else {
60                        return Err(InputError::new(input, 0, "expected mask before write"));
61                    };
62                    writes.push(Write {
63                        address,
64                        value,
65                        ones,
66                        floating,
67                    });
68                }
69            }
70        }
71
72        Ok(Self { writes })
73    }
74
75    #[must_use]
76    pub fn part1(&self) -> u64 {
77        let mut seen = FastSet::with_capacity(self.writes.len());
78        let mut total = 0;
79        for write in self.writes.iter().rev() {
80            if seen.insert(write.address) {
81                total += (write.value & write.floating) | write.ones;
82            }
83        }
84        total
85    }
86
87    #[must_use]
88    pub fn part2(&self) -> u64 {
89        let sets = self
90            .writes
91            .iter()
92            .map(|write| AddressSet {
93                ones: (write.address | write.ones) & !write.floating,
94                floating: write.floating,
95            })
96            .collect::<Vec<_>>();
97
98        // Overlapping writes must share an address, so index later writes by the low 8 bits of
99        // their addresses to avoid checking every pair
100        let width = sets.len().div_ceil(64);
101        let mut index = vec![0_u64; (1 << INDEX_BITS) * width];
102        let mut candidates = vec![0_u64; width];
103        let mut overwritten = Vec::new();
104        let mut total = 0;
105        for (i, &set) in sets.iter().enumerate().rev() {
106            // For each index key, read the candidates from the index, then add this set
107            candidates.fill(0);
108            for key in set.index_keys() {
109                let indexed = &index[key * width..][..width];
110                for (candidate, &word) in candidates.iter_mut().zip(indexed) {
111                    *candidate |= word;
112                }
113                index[key * width + (i / 64)] |= 1 << (i % 64);
114            }
115
116            overwritten.clear();
117            for (w, &word) in candidates.iter().enumerate() {
118                for (bit, _) in BitIterator::ones(word) {
119                    let later = w * 64 + bit as usize;
120                    if let Some(intersection) = set.intersect(sets[later]) {
121                        overwritten.push(intersection);
122                    }
123                }
124            }
125            total += set.size_excluding(&overwritten) * self.writes[i].value;
126        }
127
128        total
129    }
130}
131
132impl AddressSet {
133    #[inline]
134    fn intersect(self, other: Self) -> Option<Self> {
135        let disjoint = (self.ones ^ other.ones) & !(self.floating | other.floating);
136        (disjoint == 0).then_some(Self {
137            ones: self.ones | other.ones,
138            floating: self.floating & other.floating,
139        })
140    }
141
142    #[inline]
143    fn size(self) -> u64 {
144        1 << self.floating.count_ones()
145    }
146
147    #[inline]
148    fn index_keys(self) -> impl Iterator<Item = usize> {
149        let floating = (self.floating & INDEX_MASK) as usize;
150        let ones = (self.ones & INDEX_MASK) as usize;
151        std::iter::successors(Some(floating), move |&subset| {
152            (subset != 0).then(|| (subset - 1) & floating)
153        })
154        .map(move |subset| ones | subset)
155    }
156
157    #[inline]
158    fn size_excluding(self, others: &[AddressSet]) -> u64 {
159        let mut size = self.size();
160        for (i, &other) in others.iter().enumerate() {
161            if let Some(intersection) = self.intersect(other) {
162                size -= intersection.size_excluding(&others[(i + 1)..]);
163            }
164        }
165        size
166    }
167}
168
169examples!(Day14 -> (u64, u64) [
170    {
171        input: "mask = XXXXXXXXXXXXXXXXXXXXXXXXXXXXX1XXXX0X\n\
172            mem[8] = 11\n\
173            mem[7] = 101\n\
174            mem[8] = 0",
175        part1: 165
176    },
177    {
178        input: "mask = 000000000000000000000000000000X1001X\n\
179            mem[42] = 100\n\
180            mask = 00000000000000000000000000000000X0XX\n\
181            mem[26] = 1",
182        part2: 208
183    },
184]);