; Stolen with love from Haskell's Data.Map. (define-library (csc map) (export delete difference empty empty? insert intersect list->map lookup map->list map? not-found-error? singleton size subset? union union-with) (import (scheme base) (only (scheme case-lambda) case-lambda)) (begin (define-record-type (make-map cmp root) map? (cmp map-cmp) (root map-root)) (define-record-type (make-node size key val left right) node? (size node-size) (key node-key) (val node-val) (left node-left) (right node-right)) (define-syntax match (syntax-rules (node tip) ((match n (tip case1 ...) ((node k x l r) case2 ...)) (if n (let ((k (node-key n)) (x (node-val n)) (l (node-left n)) (r (node-right n))) case2 ...) (begin case1 ...))) ((match n ((node k x l r) body ...)) (match n (tip (error "no match")) ((node k x l r) body ...))))) (define (empty cmp) (make-map cmp #f)) (define (empty? m) (not (map-root m))) (define (singleton-node key val) (make-node 1 key val #f #f)) (define (singleton cmp key val) (make-map cmp (singleton-node key val))) (define (sz n) (if n (node-size n) 0)) (define (size m) (sz (map-root m))) (define-record-type (make-not-found-error) not-found-error?) (define *not-found-error* (make-not-found-error)) (define lookup (case-lambda ((m k) (let loop ((n (map-root m))) (unless n (raise *not-found-error*)) (let ((ord ((map-cmp m) k (node-key n)))) (cond ((negative? ord) (loop (node-left n))) ((positive? ord) (loop (node-right n))) (else (node-val n)))))) ((m k def) (guard (err ((not-found-error? err) def)) (lookup m k))))) ; The bin constructor maintains the size of the tree. (define (bin k x l r) (make-node (+ (sz l) (sz r) 1) k x l r)) ; https://hackage.haskell.org/package/containers-0.4.0.0/docs/src/Data-Map.html#delta (define delta 4) (define ratio 2) ; This is where the magic happens... (define (single-l k1 x1 t1 n) (match n ((node k2 x2 t2 t3) (bin k2 x2 (bin k1 x1 t1 t2) t3)))) (define (single-r k1 x1 n t3) (match n ((node k2 x2 t1 t2) (bin k2 x2 t1 (bin k1 x1 t2 t3))))) (define (double-l k1 x1 t1 n1) (match n1 ((node k2 x2 n2 t4) (match n2 ((node k3 x3 t2 t3) (bin k3 x3 (bin k1 x1 t1 t2) (bin k2 x2 t3 t4))))))) (define (double-r k1 x1 n1 t4) (match n1 ((node k2 x2 t1 n2) (match n2 ((node k3 x3 t2 t3) (bin k3 x3 (bin k2 x2 t1 t2) (bin k1 x1 t3 t4))))))) (define (rotate-l k x l r) (if (< (sz (node-left r)) (* ratio (sz (node-right r)))) (single-l k x l r) (double-l k x l r))) (define (rotate-r k x l r) (if (< (sz (node-right l)) (* ratio (sz (node-left l)))) (single-r k x l r) (double-r k x l r))) (define (balance k x l r) (define size-l (sz l)) (define size-r (sz r)) (cond ((<= (+ size-l size-r) 1) (bin k x l r)) ((> size-r (* delta size-l)) (rotate-l k x l r)) ((> size-l (* delta size-r)) (rotate-r k x l r)) (else (bin k x l r)))) (define (insert-max t kx x) (match t (tip (singleton-node kx x)) ((node ky y l r) (balance ky y l (insert-max r kx x))))) (define (insert-min t kx x) (match t (tip (singleton-node kx x)) ((node ky y l r) (balance ky y (insert-min l kx x) r)))) ; This is the "general purpose map function". All other map operations are ; defined in terms of join. ; https://www.cs.cmu.edu/~guyb/papers/BFS16.pdf (define (join kx x l r) (cond ((not l) (insert-min r kx x)) ((not r) (insert-max l kx x)) (else (match l ((node ky y ly ry) (define size-l (node-size l)) (match r ((node kz z lz rz) (define size-r (node-size r)) (cond ((<= (* delta size-l) size-r) (balance kz z (join kx x l lz) rz)) ((<= (* delta size-r) size-l) (balance ky y ly (join kx x ry r))) (else (bin kx x l r)))))))))) (define (split cmp t k) (match t (tip (values #f *not-found-error* #f)) ((node km m l r) (define ord (cmp k km)) (cond ((zero? ord) (values l m r)) ((negative? ord) (let-values (((ll b lr) (split cmp l k))) (values ll b (join km m lr r)))) (else (let-values (((rl b rr) (split cmp r k))) (values (join km m l rl) b rr))))))) (define (split-last t) (match t ((node kx x l r) (if r (let-values (((t2 kx2 x2) (split-last r))) (values (join kx x l t2) kx2 x2)) (values l kx x))))) (define (join2 tl tr) (if tl (let-values (((tl2 kx x) (split-last tl))) (join kx x tl2 tr)) tr)) (define (insert t kx x) (define cmp (map-cmp t)) (define-values (tl m tr) (split cmp (map-root t) kx)) (make-map cmp (join kx x tl tr))) (define (delete t k) (define cmp (map-cmp t)) (define-values (tl m tr) (split cmp (map-root t) k)) (make-map cmp (join2 tl tr))) (define (foldl f acc l) (if (null? l) acc (foldl f (f acc (car l)) (cdr l)))) (define (union-node cmp f t1 t2) (cond ((not t2) t1) ((not t1) t2) (else (match t1 ((node k1 x1 l1 r1) (define-values (l2 x2 r2) (split cmp t2 k1)) (define tl (union-node cmp f l1 l2)) (define tr (union-node cmp f r1 r2)) (define x (if (eq? x2 *not-found-error*) x1 (f x1 x2))) (join k1 x tl tr)))))) (define (union-with f m . ms) (define cmp (map-cmp m)) (make-map cmp (foldl (lambda (acc x) (union-node cmp f acc (map-root x))) (map-root m) ms))) (define (union m . ms) (apply union-with (lambda (l r) r) m ms)) (define (intersect-node cmp t1 t2) (cond ((not (and t1 t2)) #f) (else (match t2 ((node k2 x2 l2 r2) (define-values (l1 b r1) (split cmp t1 k2)) (define tl (intersect-node cmp l1 l2)) (define tr (intersect-node cmp r1 r2)) (if (eq? b *not-found-error*) (join2 tl tr) (join k2 x2 tl tr))))))) (define (intersect m . ms) (define cmp (map-cmp m)) (make-map cmp (foldl (lambda (acc x) (intersect-node cmp acc (map-root x))) (map-root m) ms))) (define (difference-node cmp t1 t2) (and t1 (match t2 (tip t1) ((node k2 x2 l2 r2) (define-values (l1 b r1) (split cmp t1 k2)) (define tl (difference-node cmp l1 l2)) (define tr (difference-node cmp r1 r2)) (join2 tl tr))))) (define (difference m . ms) (define cmp (map-cmp m)) (make-map cmp (foldl (lambda (acc x) (difference-node cmp acc (map-root x))) (map-root m) ms))) (define (subset-node cmp a b) (match a (tip #t) ((node k x al ar) (define-values (bl found? br) (split cmp b k)) (and (not (eq? found? *not-found-error*)) (subset-node cmp al bl) (subset-node cmp ar br))))) (define (subset? a b) (define cmp (map-cmp a)) (make-map cmp (subset-node cmp (map-root a) (map-root b)))) (define (list->map cmp l) (foldl (lambda (acc x) (insert acc (car x) (cdr x))) (empty cmp) l)) (define (map->list m) (let loop ((n (map-root m)) (acc '())) (match n (tip acc) ((node k v l r) (loop l (cons (cons k v) (loop r acc)))))))))