(define-library (csc hash-map) (export alist->map hash-bytevector insert key-not-found-error? lookup make-map delete map->alist map-for-each map? merge) (import (scheme base) (only (csc match) define-match-record-type match) (only (scheme case-lambda) case-lambda)) (begin (define-record-type (make-key-hash k hash) key-hash? (k key-hash-value) (hash key-hash-hash)) (define (cmp-key-hash k1 k2 cmp) (define d (- (key-hash-hash k2) (key-hash-hash k1))) (if (zero? d) (cmp (key-hash-value k1) (key-hash-value k2)) d)) (define-match-record-type (make-node color key-hash val left right) node? %node (color node-color) (key-hash node-key) (val node-value) (left node-left) (right node-right)) (define (red? n) (and (not (null? n)) (eq? 'red (node-color n)))) (define (black? n) (or (null? n) (eq? 'black (node-color n)))) (define (node-kv n) (cons (node-key n) (node-value n))) (define (node-colored color kv left right) (make-node color (car kv) (cdr kv) left right)) (define (red-node kv left right) (node-colored 'red kv left right)) (define (black-node kv left right) (node-colored 'black kv left right)) (define (rebalance-left m) (define p (node-left m)) (define u (node-right m)) (cond ((or (and (red? p) (red? (node-left p)) (red? u)) (and (red? p) (red? (node-right p)) (red? u))) ; b r ; / \ / \ ; r r => b b ; / / ; r r ; b r ; / \ / \ ; r r => b b ; \ \ ; r r (red-node (node-kv m) (black-node (node-kv p) (node-left p) (node-right p)) (black-node (node-kv u) (node-left u) (node-right u)))) ((and (red? p) (red? (node-right p)) (black? u)) ; b b ; / \ / \ ; r b => r r ; \ \ ; r b (let ((n (node-right p))) (make-node 'black (node-key n) (node-value n) (make-node 'red (node-key p) (node-value p) (node-left p) (node-left n)) (make-node 'red (node-key m) (node-value m) (node-right n) u)))) ((and (red? p) (red? (node-left p)) (black? u)) ; b b ; / \ / \ ; r b => r r ; / \ ; r b (make-node 'black (node-key p) (node-value p) (node-left p) (make-node 'red (node-key m) (node-value m) (node-right p) u))) (else m))) (define (rebalance-right m) (define u (node-left m)) (define p (node-right m)) (cond ((or (and (red? u) (red? p) (red? (node-left p))) (and (red? u) (red? p) (red? (node-right p)))) ; b r ; / \ / \ ; r r => b b ; \ \ ; r r ; b r ; / \ / \ ; r r => b b ; / / ; r r (make-node 'red (node-key m) (node-value m) (make-node 'black (node-key u) (node-value u) (node-left u) (node-right u)) (make-node 'black (node-key p) (node-value p) (node-left p) (node-right p)))) ((and (black? u) (red? p) (red? (node-left p))) ; b b ; / \ / \ ; b r => r r ; / / ; r b (let ((n (node-left p))) (make-node 'black (node-key n) (node-value n) (make-node 'red (node-key m) (node-value m) u (node-left n)) (make-node 'red (node-key p) (node-value p) (node-right n) (node-right p))))) ((and (black? u) (red? p) (red? (node-right p))) ; b b ; / \ / \ ; b r => r r ; \ / ; r b (make-node 'black (node-key p) (node-value p) (make-node 'red (node-key m) (node-value m) u (node-left p)) (node-right p))) (else m))) (define (insert-node n k v cmp) (match n ('() (make-node 'red k v '() '())) ((% %node color node-key node-val left right) (define ord (cmp-key-hash k node-key cmp)) (cond ((negative? ord) (rebalance-left (make-node color node-key node-val (insert-node left k v cmp) right))) ((zero? ord) (make-node color k v left right)) (else (rebalance-right (make-node color node-key node-val left (insert-node right k v cmp)))))))) (define-record-type (construct-map hash cmp root) map? (hash map-raw-hash) (cmp map-cmp) (root map-root)) (define (make-map hash cmp) (construct-map hash cmp '())) (define (shuffle n) (truncate-remainder (* #x9e3779b97f4a7c55 n) #x10000000000000000)) (define (map-hash m) (lambda (k) (shuffle ((map-raw-hash m) k)))) (define (insert m k v) (define res (insert-node (map-root m) (make-key-hash k ((map-hash m) k)) v (map-cmp m))) (construct-map (map-raw-hash m) (map-cmp m) (make-node 'black (node-key res) (node-value res) (node-left res) (node-right res)))) (define-record-type (make-key-not-found-error) key-not-found-error?) (define *key-not-found-error* (make-key-not-found-error)) (define lookup (case-lambda ((m k) (define cmp (map-cmp m)) (define k* (make-key-hash k ((map-hash m) k))) (let loop ((n (map-root m))) (match n ('() (raise *key-not-found-error*)) ((% %node _ node-key node-val left right) (define ord (cmp-key-hash k* node-key cmp)) (cond ((negative? ord) (loop left)) ((zero? ord) node-val) (else (loop right))))))) ((m k def) (guard (e ((key-not-found-error? e) def)) (lookup m k))))) (define (map-for-each f m) (let loop ((n (map-root m))) (match n ((% %node _ node-key node-val left right) (loop left) (f (key-hash-value node-key) node-val) (loop right))))) (define (map->alist m) (let ((alist '())) (map-for-each (lambda (k v) (set! alist (cons (cons k v) alist))) m) alist)) (define (alist->map hash cmp alist) (let loop ((alist alist) (m (make-map hash cmp))) (if (null? alist) m (loop (cdr alist) (insert m (caar alist) (cdar alist)))))) (define (hash-bytevector b) (let loop ((i 0) (hash 0)) (if (>= i (bytevector-length b)) hash (loop (+ 1 i) (+ (* hash #x100) (bytevector-u8-ref b i)))))) (define (merge2 m1 m2) (let ((m1 m1)) (map-for-each (lambda (k v) (set! m1 (insert m1 k v))) m2) m1)) (define (merge m . m*) (let loop ((m* m*) (m m)) (match m* ('() m) ((head . tail) (loop tail (merge2 m head)))))) ; Jinkies! (define (delete m k) (define cmp (map-cmp m)) (define k* (make-key-hash k ((map-hash m) k))) ; local variables (define need-fix #f) (define replacement-node #f) (define (fix-black-height-left p) ; n.b.: s must not be nil, because we deleted a black node and so ; there must be at least one node in s to balance out the ; black height. (define s (node-right p)) (define n (node-left p)) (define c (node-left s)) (define d (node-right s)) (set! need-fix #f) (cond ((red? s) ; p s ; / \ / \ ; n s => p d ; / \ / \ ; c d n c (black-node (node-kv s) (fix-black-height-left (red-node (node-kv p) n c)) d)) ((red? d) (node-colored (node-color p) (node-kv s) (black-node (node-kv p) n c) (black-node (node-kv d) (node-left d) (node-right d)))) ((red? c) (fix-black-height-left (node-colored (node-color p) (node-kv p) n (black-node (node-kv c) (node-left c) (red-node (node-kv s) (node-right c) d))))) ((red? p) (black-node (node-kv p) n (red-node (node-kv s) c d))) (else (set! need-fix #t) ; This is the only recursive case. (black-node (node-kv p) n (red-node (node-kv s) c d))))) (define (fix-black-height-right p) (define s (node-left p)) (define n (node-right p)) (define c (node-right s)) (define d (node-left s)) (set! need-fix #f) (cond ((red? s) (black-node (node-kv s) d (fix-black-height-right (red-node (node-kv p) c n)))) ((red? d) (node-colored (node-color p) (node-kv s) (black-node (node-kv d) (node-left d) (node-right d)) (black-node (node-kv p) c n))) ((red? c) ; p p ; / \ / \ ; s n => c n ; / \ / ; d c s ; / ; d (fix-black-height-right (node-colored (node-color p) (node-kv p) (black-node (node-kv c) (red-node (node-kv s) d (node-left c)) (node-right c)) n))) ((red? p) (black-node (node-kv p) (red-node (node-kv s) d c) n)) (else (set! need-fix #t) (black-node (node-kv p) (red-node (node-kv s) d c) n)))) (construct-map (map-raw-hash m) cmp (let loop ((n (map-root m))) (match n ('() '()) ((% %node color nk nv left right) (define ord (cmp-key-hash k* nk cmp)) (cond ((negative? ord) (let* ((left* (loop left)) (n* (make-node color nk nv left* right))) (if need-fix (fix-black-height-left n*) n*))) ((positive? ord) (let* ((right* (loop right)) (n* (make-node color nk nv left right*))) (if need-fix (fix-black-height-right n*) n*))) ((and (not (null? left)) (not (null? right))) (let* ((left* (let find-max ((r left)) (match r ((% %node _ _ _ r-left '()) (set! replacement-node r) (cond ((red? r) '()) ((null? r-left) (set! need-fix #t) '()) (else (black-node (node-kv r-left) (node-left r-left) (node-right r-left))))) ((% %node r-color r-k r-v r-left r-right) (define n* (make-node r-color r-k r-v r-left (find-max r-right))) (if need-fix (fix-black-height-right n*) n*))))) (n* (node-colored color (node-kv replacement-node) left* right))) (if need-fix (fix-black-height-left n*) n*))) ((red? n) '()) ((and (null? left) (null? right)) (set! need-fix #t) '()) (else (let ((child (if (null? left) right left))) (black-node (node-kv child) (node-left child) (node-right child))))))))))))