iter: :when, :while, :let clause extensions for all for macros
ober
a4078288451578195e9e75cdc66eaa5f56d4457e
--- a/lib/std/iter.sls +++ b/lib/std/iter.sls @@ -5,6 +5,11 @@ ;;; with iterator constructors: in-list, in-vector, in-range, ;;; in-string, in-hash-keys, in-hash-values, in-hash-pairs, ;;; in-naturals, in-indexed +;;; +;;; Clause extensions (Clojure-style): +;;; :when expr — skip iteration when expr is #f +;;; :while expr — stop iteration when expr is #f +;;; :let ((var expr) ...) — bind intermediate values (library (std iter) (export @@ -20,7 +25,6 @@ (jerboa runtime)) ;; Iterator constructors — return plain lists for simplicity - ;; (Gerbil iterators are more complex, but lists suffice for porting) (define (in-list lst) lst) @@ -53,85 +57,248 @@ (case-lambda (() (in-naturals 0)) ((start) - ;; Returns an infinite-ish list — but for/collect with zip will stop - ;; at the shorter list. Use iota for bounded ranges. - ;; For practical use, generate up to a reasonable limit. - ;; In real Gerbil this is lazy; here we rely on for macros to limit. (let loop ([i start] [acc '()] [n 0]) (if (>= n 100000) (reverse acc) (loop (+ i 1) (cons i acc) (+ n 1))))))) (define (in-indexed lst) - ;; Returns list of (index . element) pairs (let loop ([rest lst] [i 0] [acc '()]) (if (null? rest) (reverse acc) (loop (cdr rest) (+ i 1) (cons (cons i (car rest)) acc))))) + ;; ========================================================================= + ;; Clause-aware for macros + ;; ========================================================================= + + ;; Shared expand-time helpers, duplicated in each macro's (let ...) + ;; to ensure they're at the correct phase. + ;; + ;; kw?: check if syntax s is a Jerboa keyword with given name + ;; e.g., (kw? #'x "when") checks if x is the keyword when: + ;; At syntax level, when: has datum symbol "when:" (NOT "#:when") + ;; + ;; binding-id?: check if syntax is a non-keyword identifier + + ;; Internal helper macro for general clause expansion. + ;; Used by for/collect, for, for/fold, for/or, for/and. + (define-syntax %clause-expand + (let () + (define (kw? s name) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (and (symbol? d) + (string=? (symbol->string d) + (string-append name ":")))))) + + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + + (define (split-mods clauses) + (syntax-case clauses () + [() (values '() #'())] + [(k expr . rest) + (or (kw? #'k "when") (kw? #'k "while")) + (let-values ([(ms remaining) (split-mods #'rest)]) + (values (cons (list #'k #'expr) ms) remaining))] + [(k binds . rest) + (kw? #'k "let") + (let-values ([(ms remaining) (split-mods #'rest)]) + (values (cons (list #'k #'binds) ms) remaining))] + [other (values '() #'other)])) + + (define (any-while? mods) + (and (pair? mods) + (or (kw? (caar mods) "while") + (any-while? (cdr mods))))) + + (define (wrap-mods mods inner stop-sym) + (if (null? mods) inner + (let ([k (caar mods)] [arg (cadar mods)] [rest (cdr mods)]) + (cond + [(kw? k "when") + (with-syntax ([w (wrap-mods rest inner stop-sym)] [e arg]) + #'(when e w))] + [(kw? k "while") + (with-syntax ([w (wrap-mods rest inner stop-sym)] + [e arg] [stop stop-sym]) + #'(if e w (set! stop #t)))] + [(kw? k "let") + (with-syntax ([w (wrap-mods rest inner stop-sym)] [b arg]) + #'(let b w))])))) + + (define (expand-clauses clauses leaf) + (syntax-case clauses () + [() leaf] + [((var iter-expr) . after) + (binding-id? #'var) + (let-values ([(mods remaining) (split-mods #'after)]) + (let ([inner (expand-clauses remaining leaf)]) + (if (any-while? mods) + (with-syntax ([wrapped (wrap-mods mods inner + (datum->syntax #'var '%stop?))] + [%stop? (datum->syntax #'var '%stop?)]) + #'(let loop ([lst iter-expr]) + (when (pair? lst) + (let ([var (car lst)] [%stop? #f]) + wrapped + (unless %stop? (loop (cdr lst))))))) + (with-syntax ([wrapped (wrap-mods mods inner #f)]) + #'(let loop ([lst iter-expr]) + (when (pair? lst) + (let ([var (car lst)]) + wrapped) + (loop (cdr lst))))))))])) + + (lambda (stx) + (syntax-case stx () + [(_ (clause ...) leaf-expr) + (expand-clauses #'(clause ...) #'leaf-expr)])))) + ;; for — side-effecting iteration (define-syntax for - (syntax-rules () - [(_ ((var iter-expr)) body ...) - (for-each (lambda (var) body ...) iter-expr)] - [(_ ((var1 iter1) (var2 iter2)) body ...) - (let loop ([l1 iter1] [l2 iter2]) - (when (and (pair? l1) (pair? l2)) - (let ([var1 (car l1)] [var2 (car l2)]) - body ... - (loop (cdr l1) (cdr l2)))))] - [(_ ((var1 iter1) (var2 iter2) (var3 iter3)) body ...) - (let loop ([l1 iter1] [l2 iter2] [l3 iter3]) - (when (and (pair? l1) (pair? l2) (pair? l3)) - (let ([var1 (car l1)] [var2 (car l2)] [var3 (car l3)]) - body ... - (loop (cdr l1) (cdr l2) (cdr l3)))))])) + (let () + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + (lambda (stx) + (syntax-case stx () + [(_ ((var iter-expr)) body ...) + (binding-id? #'var) + #'(for-each (lambda (var) body ...) iter-expr)] + [(_ ((var1 iter1) (var2 iter2)) body ...) + (and (binding-id? #'var1) (binding-id? #'var2)) + #'(let loop ([l1 iter1] [l2 iter2]) + (when (and (pair? l1) (pair? l2)) + (let ([var1 (car l1)] [var2 (car l2)]) + body ... + (loop (cdr l1) (cdr l2)))))] + [(_ ((var1 iter1) (var2 iter2) (var3 iter3)) body ...) + (and (binding-id? #'var1) (binding-id? #'var2) (binding-id? #'var3)) + #'(let loop ([l1 iter1] [l2 iter2] [l3 iter3]) + (when (and (pair? l1) (pair? l2) (pair? l3)) + (let ([var1 (car l1)] [var2 (car l2)] [var3 (car l3)]) + body ... + (loop (cdr l1) (cdr l2) (cdr l3)))))] + [(_ (clause ...) body ...) + #'(%clause-expand (clause ...) (begin body ...))])))) ;; for/collect — collect results into a list (define-syntax for/collect - (syntax-rules () - [(_ ((var iter-expr)) body ...) - (map (lambda (var) body ...) iter-expr)] - [(_ ((var1 iter1) (var2 iter2)) body ...) - (let loop ([l1 iter1] [l2 iter2] [acc '()]) - (if (or (null? l1) (null? l2)) - (reverse acc) - (let ([var1 (car l1)] [var2 (car l2)]) - (loop (cdr l1) (cdr l2) (cons (begin body ...) acc)))))])) + (let () + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + (lambda (stx) + (syntax-case stx () + [(_ ((var iter-expr)) body ...) + (binding-id? #'var) + #'(map (lambda (var) body ...) iter-expr)] + [(_ ((var1 iter1) (var2 iter2)) body ...) + (and (binding-id? #'var1) (binding-id? #'var2)) + #'(let loop ([l1 iter1] [l2 iter2] [acc '()]) + (if (or (null? l1) (null? l2)) + (reverse acc) + (let ([var1 (car l1)] [var2 (car l2)]) + (loop (cdr l1) (cdr l2) (cons (begin body ...) acc)))))] + [(_ (clause ...) body ...) + (with-syntax ([%acc (datum->syntax (car (syntax->list stx)) '%acc)]) + #`(let ([%acc '()]) + (%clause-expand (clause ...) (set! %acc (cons (begin body ...) %acc))) + (reverse %acc)))])))) ;; for/fold — fold with accumulator (define-syntax for/fold - (syntax-rules () - [(_ ((acc init)) ((var iter-expr)) body ...) - (let loop ([rest iter-expr] [acc init]) - (if (null? rest) acc - (let ([var (car rest)]) - (loop (cdr rest) (begin body ...)))))] - [(_ ((acc init)) ((var1 iter1) (var2 iter2)) body ...) - (let loop ([l1 iter1] [l2 iter2] [acc init]) - (if (or (null? l1) (null? l2)) acc - (let ([var1 (car l1)] [var2 (car l2)]) - (loop (cdr l1) (cdr l2) (begin body ...)))))])) + (let () + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + (lambda (stx) + (syntax-case stx () + [(_ ((acc init)) ((var iter-expr)) body ...) + (binding-id? #'var) + #'(let loop ([rest iter-expr] [acc init]) + (if (null? rest) acc + (let ([var (car rest)]) + (loop (cdr rest) (begin body ...)))))] + [(_ ((acc init)) ((var1 iter1) (var2 iter2)) body ...) + (and (binding-id? #'var1) (binding-id? #'var2)) + #'(let loop ([l1 iter1] [l2 iter2] [acc init]) + (if (or (null? l1) (null? l2)) acc + (let ([var1 (car l1)] [var2 (car l2)]) + (loop (cdr l1) (cdr l2) (begin body ...)))))] + [(_ ((acc init)) (clause ...) body ...) + #`(let ([acc init]) + (%clause-expand (clause ...) (set! acc (begin body ...))) + acc)])))) ;; for/or — return first truthy result (define-syntax for/or - (syntax-rules () - [(_ ((var iter-expr)) body ...) - (let loop ([rest iter-expr]) - (if (null? rest) #f - (let ([var (car rest)]) - (or (begin body ...) (loop (cdr rest))))))])) + (let () + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + (lambda (stx) + (syntax-case stx () + [(_ ((var iter-expr)) body ...) + (binding-id? #'var) + #'(let loop ([rest iter-expr]) + (if (null? rest) #f + (let ([var (car rest)]) + (or (begin body ...) (loop (cdr rest))))))] + [(_ (clause ...) body ...) + #'(call/cc (lambda (return) + (%clause-expand (clause ...) + (let ([%result (begin body ...)]) + (when %result (return %result)))) + #f))])))) ;; for/and — return #f if any result is #f (define-syntax for/and - (syntax-rules () - [(_ ((var iter-expr)) body ...) - (let loop ([rest iter-expr]) - (if (null? rest) #t - (let ([var (car rest)]) - (and (begin body ...) (loop (cdr rest))))))])) + (let () + (define (binding-id? s) + (and (identifier? s) + (let ([d (syntax->datum s)]) + (or (not (symbol? d)) + (let ([str (symbol->string d)]) + (or (= (string-length str) 0) + (not (char=? (string-ref str (- (string-length str) 1)) #\:)))))))) + (lambda (stx) + (syntax-case stx () + [(_ ((var iter-expr)) body ...) + (binding-id? #'var) + #'(let loop ([rest iter-expr]) + (if (null? rest) #t + (let ([var (car rest)]) + (and (begin body ...) (loop (cdr rest))))))] + [(_ (clause ...) body ...) + #'(call/cc (lambda (return) + (%clause-expand (clause ...) + (unless (begin body ...) (return #f))) + #t))])))) ;; ========== better2 #7: I/O iterators ========== - ;; Read all datums from a port using read (define in-port (case-lambda [() (in-port (current-input-port))] @@ -143,7 +310,6 @@ (reverse acc) (loop (cons datum acc)))))])) - ;; Read all lines from a port (define in-lines (case-lambda [() (in-lines (current-input-port))] @@ -154,7 +320,6 @@ (reverse acc) (loop (cons line acc)))))])) - ;; Read all characters from a port (define in-chars (case-lambda [() (in-chars (current-input-port))] @@ -165,7 +330,6 @@ (reverse acc) (loop (cons ch acc)))))])) - ;; Read all bytes from a binary port (define in-bytes (case-lambda [() (in-bytes (current-input-port))] @@ -176,7 +340,6 @@ (reverse acc) (loop (cons b acc)))))])) - ;; Iterate over results of a thunk until it returns eof-object (define (in-producer thunk . sentinel) (let ([stop? (if (null? sentinel) eof-object? @@ -188,4 +351,4 @@ (reverse acc) (loop (cons val acc))))))) - ) ;; end library +) ;; end library new file mode 100644 --- /dev/null +++ b/tests/test-for-clauses.ss @@ -0,0 +1,164 @@ +(import (jerboa prelude)) + +(def test-count 0) +(def pass-count 0) + +(defrule (test name body ...) + (begin + (set! test-count (+ test-count 1)) + (guard (exn [#t + (displayln (str "FAIL: " name)) + (displayln (str " Error: " (if (message-condition? exn) + (condition-message exn) exn)))]) + body ... + (set! pass-count (+ pass-count 1)) + (displayln (str "PASS: " name))))) + +(defrule (assert-equal got expected msg) + (unless (equal? got expected) + (error 'assert msg (list 'got: got 'expected: expected)))) + +(defrule (assert-true val msg) + (unless val (error 'assert msg))) + +;; ========================================================================= +;; Backwards compatibility — existing syntax still works +;; ========================================================================= + +(test "for/collect single binding" + (assert-equal (for/collect ([x (in-range 5)]) (* x x)) + '(0 1 4 9 16) "squares")) + +(test "for/collect two bindings" + (assert-equal (for/collect ([x '(1 2)] [y '(a b)]) (list x y)) + '((1 a) (2 b)) "zipped")) + +(test "for single binding side effect" + (let ([acc '()]) + (for ([x '(1 2 3)]) (set! acc (cons x acc))) + (assert-equal (reverse acc) '(1 2 3) "side effects"))) + +(test "for/fold single binding" + (assert-equal (for/fold ([sum 0]) ([x (in-range 5)]) (+ sum x)) + 10 "sum 0-4")) + +(test "for/or single binding" + (assert-equal (for/or ([x '(1 3 4 7)]) (and (even? x) x)) + 4 "first even")) + +(test "for/and single binding" + (assert-true (for/and ([x '(2 4 6)]) (even? x)) + "all even")) + +;; ========================================================================= +;; :when clause +;; ========================================================================= + +(test "for/collect :when" + (assert-equal (for/collect ([x (in-range 10)] when: (even? x)) x) + '(0 2 4 6 8) "even only")) + +(test "for/collect :when with body" + (assert-equal (for/collect ([x (in-range 10)] when: (> x 5)) (* x 10)) + '(60 70 80 90) "filtered and transformed")) + +(test "for :when side effect" + (let ([acc '()]) + (for ([x (in-range 10)] when: (even? x)) + (set! acc (cons x acc))) + (assert-equal (reverse acc) '(0 2 4 6 8) "side effects with when"))) + +(test "for/fold :when" + (assert-equal (for/fold ([sum 0]) ([x (in-range 10)] when: (even? x)) + (+ sum x)) + 20 "sum of evens 0-8")) + +(test "for/or :when" + (assert-equal (for/or ([x '(1 3 5 6 7)] when: (even? x)) x) + 6 "first match after filter")) + +(test "for/and :when" + (assert-true (for/and ([x '(1 2 3 4 5 6)] when: (even? x)) (< x 10)) + "all filtered elements < 10")) + +;; ========================================================================= +;; :while clause +;; ========================================================================= + +(test "for/collect :while" + (assert-equal (for/collect ([x (in-range 10)] while: (< x 5)) x) + '(0 1 2 3 4) "take while < 5")) + +(test "for/collect :while stops early" + (assert-equal (for/collect ([x '(1 2 3 10 4 5)] while: (< x 10)) x) + '(1 2 3) "stops at 10")) + +(test "for/fold :while" + (assert-equal (for/fold ([sum 0]) ([x (in-range 100)] while: (< x 5)) + (+ sum x)) + 10 "sum while < 5")) + +;; ========================================================================= +;; :let clause +;; ========================================================================= + +(test "for/collect :let" + (assert-equal (for/collect ([x (in-range 5)] + let: ([y (* x x)]) + when: (even? y)) + y) + '(0 4 16) "let + when")) + +(test "for/collect :let multiple bindings" + (assert-equal (for/collect ([x (in-range 1 4)] + let: ([y (* x 10)] [z (+ x 1)])) + (list x y z)) + '((1 10 2) (2 20 3) (3 30 4)) "multi-let")) + +;; ========================================================================= +;; Nested bindings (Clojure for comprehension) +;; ========================================================================= + +(test "for/collect two bindings zips" + (assert-equal (for/collect ([x '(1 2)] [y '(a b)]) (list x y)) + '((1 a) (2 b)) "zip, not cross-product")) + +(test "for/collect general path cross-product" + (assert-equal (for/collect ([x '(1 2)] when: #t [y '(a b)]) (list x y)) + '((1 a) (1 b) (2 a) (2 b)) "cross-product via clauses")) + +(test "for/collect nested with :when" + (assert-equal (for/collect ([x (in-range 1 4)] + [y (in-range 1 4)] + when: (not (= x y))) + (list x y)) + '((1 2) (1 3) (2 1) (2 3) (3 1) (3 2)) "permutations")) + +;; ========================================================================= +;; Combined clauses +;; ========================================================================= + +(test "for/collect :when + :while" + (assert-equal (for/collect ([x (in-range 20)] + when: (even? x) + while: (< x 10)) + x) + '(0 2 4 6 8) "when + while")) + +(test "for/collect :let + :when + nested" + (assert-equal (for/collect ([x (in-range 1 5)] + let: ([sq (* x x)]) + when: (odd? sq) + [y '(10 20)]) + (+ sq y)) + '(11 21 19 29) "complex comprehension")) + +;; ========================================================================= +;; Summary +;; ========================================================================= +(newline) +(displayln (str "=========================================")) +(displayln (str "Results: " pass-count "/" test-count " passed")) +(displayln (str "=========================================")) +(when (< pass-count test-count) + (exit 1))