Source file
src/math/big/int.go
1
2
3
4
5
6
7 package big
8
9 import (
10 "fmt"
11 "io"
12 "math/rand"
13 "strings"
14 "sync"
15 )
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34 type Int struct {
35 neg bool
36 abs nat
37 }
38
39 var intOne = &Int{false, natOne}
40
41
42
43
44
45 func (x *Int) Sign() int {
46
47
48
49 if len(x.abs) == 0 {
50 return 0
51 }
52 if x.neg {
53 return -1
54 }
55 return 1
56 }
57
58
59 func (z *Int) SetInt64(x int64) *Int {
60 neg := false
61 if x < 0 {
62 neg = true
63 x = -x
64 }
65 z.abs = z.abs.setUint64(uint64(x))
66 z.neg = neg
67 return z
68 }
69
70
71 func (z *Int) SetUint64(x uint64) *Int {
72 z.abs = z.abs.setUint64(x)
73 z.neg = false
74 return z
75 }
76
77
78 func NewInt(x int64) *Int {
79
80
81 u := uint64(x)
82 if x < 0 {
83 u = -u
84 }
85 var abs []Word
86 if x == 0 {
87 } else if _W == 32 && u>>32 != 0 {
88 abs = []Word{Word(u), Word(u >> 32)}
89 } else {
90 abs = []Word{Word(u)}
91 }
92 return &Int{neg: x < 0, abs: abs}
93 }
94
95
96 func (z *Int) Set(x *Int) *Int {
97 if z != x {
98 z.abs = z.abs.set(x.abs)
99 z.neg = x.neg
100 }
101 return z
102 }
103
104
105
106
107
108
109 func (x *Int) Bits() []Word {
110
111
112
113 return x.abs
114 }
115
116
117
118
119
120
121 func (z *Int) SetBits(abs []Word) *Int {
122 z.abs = nat(abs).norm()
123 z.neg = false
124 return z
125 }
126
127
128 func (z *Int) Abs(x *Int) *Int {
129 z.Set(x)
130 z.neg = false
131 return z
132 }
133
134
135 func (z *Int) Neg(x *Int) *Int {
136 z.Set(x)
137 z.neg = len(z.abs) > 0 && !z.neg
138 return z
139 }
140
141
142 func (z *Int) Add(x, y *Int) *Int {
143 neg := x.neg
144 if x.neg == y.neg {
145
146
147 z.abs = z.abs.add(x.abs, y.abs)
148 } else {
149
150
151 if x.abs.cmp(y.abs) >= 0 {
152 z.abs = z.abs.sub(x.abs, y.abs)
153 } else {
154 neg = !neg
155 z.abs = z.abs.sub(y.abs, x.abs)
156 }
157 }
158 z.neg = len(z.abs) > 0 && neg
159 return z
160 }
161
162
163 func (z *Int) Sub(x, y *Int) *Int {
164 neg := x.neg
165 if x.neg != y.neg {
166
167
168 z.abs = z.abs.add(x.abs, y.abs)
169 } else {
170
171
172 if x.abs.cmp(y.abs) >= 0 {
173 z.abs = z.abs.sub(x.abs, y.abs)
174 } else {
175 neg = !neg
176 z.abs = z.abs.sub(y.abs, x.abs)
177 }
178 }
179 z.neg = len(z.abs) > 0 && neg
180 return z
181 }
182
183
184 func (z *Int) Mul(x, y *Int) *Int {
185 z.mul(nil, x, y)
186 return z
187 }
188
189
190
191
192 func (z *Int) mul(stk *stack, x, y *Int) {
193
194
195
196
197 if x == y {
198 z.abs = z.abs.sqr(stk, x.abs)
199 z.neg = false
200 return
201 }
202 z.abs = z.abs.mul(stk, x.abs, y.abs)
203 z.neg = len(z.abs) > 0 && x.neg != y.neg
204 }
205
206
207
208
209 func (z *Int) MulRange(a, b int64) *Int {
210 switch {
211 case a > b:
212 return z.SetInt64(1)
213 case a <= 0 && b >= 0:
214 return z.SetInt64(0)
215 }
216
217
218 neg := false
219 if a < 0 {
220 neg = (b-a)&1 == 0
221 a, b = -b, -a
222 }
223
224 z.abs = z.abs.mulRange(nil, uint64(a), uint64(b))
225 z.neg = neg
226 return z
227 }
228
229
230 func (z *Int) Binomial(n, k int64) *Int {
231 if k > n || k < 0 {
232 return z.SetInt64(0)
233 }
234
235 if k > n-k {
236 k = n - k
237 }
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259 var N, K, i, t Int
260 N.SetInt64(n)
261 K.SetInt64(k)
262 z.Set(intOne)
263 for i.Cmp(&K) < 0 {
264 z.Mul(z, t.Sub(&N, &i))
265 i.Add(&i, intOne)
266 z.Quo(z, &i)
267 }
268 return z
269 }
270
271
272
273
274 func (z *Int) Quo(x, y *Int) *Int {
275 z.abs, _ = z.abs.div(nil, nil, x.abs, y.abs)
276 z.neg = len(z.abs) > 0 && x.neg != y.neg
277 return z
278 }
279
280
281
282
283 func (z *Int) Rem(x, y *Int) *Int {
284 _, z.abs = nat(nil).div(nil, z.abs, x.abs, y.abs)
285 z.neg = len(z.abs) > 0 && x.neg
286 return z
287 }
288
289
290
291
292
293
294
295
296
297
298
299
300 func (z *Int) QuoRem(x, y, r *Int) (*Int, *Int) {
301 z.abs, r.abs = z.abs.div(nil, r.abs, x.abs, y.abs)
302 z.neg, r.neg = len(z.abs) > 0 && x.neg != y.neg, len(r.abs) > 0 && x.neg
303 return z, r
304 }
305
306
307
308
309 func (z *Int) Div(x, y *Int) *Int {
310 y_neg := y.neg
311 var r Int
312 z.QuoRem(x, y, &r)
313 if r.neg {
314 if y_neg {
315 z.Add(z, intOne)
316 } else {
317 z.Sub(z, intOne)
318 }
319 }
320 return z
321 }
322
323
324
325
326 func (z *Int) Mod(x, y *Int) *Int {
327 y0 := y
328 if z == y || alias(z.abs, y.abs) {
329 y0 = new(Int).Set(y)
330 }
331 var q Int
332 q.QuoRem(x, y, z)
333 if z.neg {
334 if y0.neg {
335 z.Sub(z, y0)
336 } else {
337 z.Add(z, y0)
338 }
339 }
340 return z
341 }
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357 func (z *Int) DivMod(x, y, m *Int) (*Int, *Int) {
358 y0 := y
359 if z == y || m == y || alias(z.abs, y.abs) || alias(m.abs, y.abs) {
360 y0 = new(Int).Set(y)
361 }
362 z.QuoRem(x, y, m)
363 if m.neg {
364 if y0.neg {
365 z.Add(z, intOne)
366 m.Sub(m, y0)
367 } else {
368 z.Sub(z, intOne)
369 m.Add(m, y0)
370 }
371 }
372 return z, m
373 }
374
375
376
377 const (
378 Trunc = ToZero
379 Floor = ToNegativeInf
380 Round = ToNearestEven
381 Ceil = ToPositiveInf
382 )
383
384
385
386
387
388
389
390
391
392
393
394 func (z *Int) Divide(x, y, r *Int, mode RoundingMode) (*Int, *Int) {
395
396 var z_abs nat
397 if z != nil {
398 z_abs = z.abs
399 }
400 var r_neg bool
401 var r_abs nat
402 if r != nil {
403 r_abs = r.abs
404 }
405 y_abs := y.abs
406 if z == y || r == y || alias(z_abs, y.abs) || alias(r_abs, y.abs) {
407 y_abs = nat(nil).set(y.abs)
408 }
409 neg := x.neg != y.neg
410 z_abs, r_abs = z_abs.div(nil, r_abs, x.abs, y.abs)
411 if len(r_abs) > 0 {
412 switch mode {
413 case Trunc:
414 r_neg = x.neg
415 case Floor:
416 r_neg = y.neg
417 if neg {
418 z_abs = z_abs.add(z_abs, natOne)
419 r_abs = r_abs.sub(y_abs, r_abs)
420 }
421 case Ceil:
422 r_neg = !y.neg
423 if !neg {
424 z_abs = z_abs.add(z_abs, natOne)
425 r_abs = r_abs.sub(y_abs, r_abs)
426 }
427 case Round:
428 switch nat(nil).mul(nil, r_abs, natTwo).cmp(y_abs) {
429 case -1:
430 r_neg = x.neg
431 case 0:
432 even := len(z_abs) == 0 || z_abs[0]&1 == 0
433 if even {
434 r_neg = x.neg
435 break
436 }
437 fallthrough
438 case 1:
439 r_neg = !x.neg
440 z_abs = z_abs.add(z_abs, natOne)
441 r_abs = r_abs.sub(y_abs, r_abs)
442 }
443 default:
444 panic("unsupported rounding mode")
445 }
446 }
447 if z != nil {
448 z.abs = z_abs
449 z.neg = neg && len(z_abs) > 0
450 }
451 if r != nil {
452 r.abs = r_abs
453 r.neg = r_neg
454 }
455 return z, r
456 }
457
458
459
460
461
462 func (x *Int) Cmp(y *Int) (r int) {
463
464
465
466
467 switch {
468 case x == y:
469
470 case x.neg == y.neg:
471 r = x.abs.cmp(y.abs)
472 if x.neg {
473 r = -r
474 }
475 case x.neg:
476 r = -1
477 default:
478 r = 1
479 }
480 return
481 }
482
483
484
485
486
487 func (x *Int) CmpAbs(y *Int) int {
488 return x.abs.cmp(y.abs)
489 }
490
491
492 func low32(x nat) uint32 {
493 if len(x) == 0 {
494 return 0
495 }
496 return uint32(x[0])
497 }
498
499
500 func low64(x nat) uint64 {
501 if len(x) == 0 {
502 return 0
503 }
504 v := uint64(x[0])
505 if _W == 32 && len(x) > 1 {
506 return uint64(x[1])<<32 | v
507 }
508 return v
509 }
510
511
512
513 func (x *Int) Int64() int64 {
514 v := int64(low64(x.abs))
515 if x.neg {
516 v = -v
517 }
518 return v
519 }
520
521
522
523 func (x *Int) Uint64() uint64 {
524 return low64(x.abs)
525 }
526
527
528 func (x *Int) IsInt64() bool {
529 if len(x.abs) <= 64/_W {
530 w := int64(low64(x.abs))
531 return w >= 0 || x.neg && w == -w
532 }
533 return false
534 }
535
536
537 func (x *Int) IsUint64() bool {
538 return !x.neg && len(x.abs) <= 64/_W
539 }
540
541
542
543 func (x *Int) Float64() (float64, Accuracy) {
544 n := x.abs.bitLen()
545 if n == 0 {
546 return 0.0, Exact
547 }
548
549
550 if n <= 53 || n < 64 && n-int(x.abs.trailingZeroBits()) <= 53 {
551 f := float64(low64(x.abs))
552 if x.neg {
553 f = -f
554 }
555 return f, Exact
556 }
557
558 return new(Float).SetInt(x).Float64()
559 }
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583 func (z *Int) SetString(s string, base int) (*Int, bool) {
584 return z.setFromScanner(strings.NewReader(s), base)
585 }
586
587
588
589 func (z *Int) setFromScanner(r io.ByteScanner, base int) (*Int, bool) {
590 if _, _, err := z.scan(r, base); err != nil {
591 return nil, false
592 }
593
594 if _, err := r.ReadByte(); err != io.EOF {
595 return nil, false
596 }
597 return z, true
598 }
599
600
601
602 func (z *Int) SetBytes(buf []byte) *Int {
603 z.abs = z.abs.setBytes(buf)
604 z.neg = false
605 return z
606 }
607
608
609
610
611 func (x *Int) Bytes() []byte {
612
613
614
615 buf := make([]byte, len(x.abs)*_S)
616 return buf[x.abs.bytes(buf):]
617 }
618
619
620
621
622
623 func (x *Int) FillBytes(buf []byte) []byte {
624
625 clear(buf)
626 x.abs.bytes(buf)
627 return buf
628 }
629
630
631
632 func (x *Int) BitLen() int {
633
634
635
636 return x.abs.bitLen()
637 }
638
639
640
641 func (x *Int) TrailingZeroBits() uint {
642 return x.abs.trailingZeroBits()
643 }
644
645
646
647
648
649
650
651 func (z *Int) Exp(x, y, m *Int) *Int {
652 return z.exp(x, y, m, false)
653 }
654
655 func (z *Int) expSlow(x, y, m *Int) *Int {
656 return z.exp(x, y, m, true)
657 }
658
659 func (z *Int) exp(x, y, m *Int, slow bool) *Int {
660
661 xWords := x.abs
662 if y.neg {
663 if m == nil || len(m.abs) == 0 {
664 return z.SetInt64(1)
665 }
666
667 inverse := new(Int).ModInverse(x, m)
668 if inverse == nil {
669 return nil
670 }
671 xWords = inverse.abs
672 }
673 yWords := y.abs
674
675 var mWords nat
676 if m != nil {
677 if z == m || alias(z.abs, m.abs) {
678 m = new(Int).Set(m)
679 }
680 mWords = m.abs
681 }
682
683 z.abs = z.abs.expNN(nil, xWords, yWords, mWords, slow)
684 z.neg = len(z.abs) > 0 && x.neg && len(yWords) > 0 && yWords[0]&1 == 1
685 if z.neg && len(mWords) > 0 {
686
687 z.abs = z.abs.sub(mWords, z.abs)
688 z.neg = false
689 }
690
691 return z
692 }
693
694
695
696
697
698
699
700
701
702
703
704
705 func (z *Int) GCD(x, y, a, b *Int) *Int {
706 if len(a.abs) == 0 || len(b.abs) == 0 {
707 lenA, lenB, negA, negB := len(a.abs), len(b.abs), a.neg, b.neg
708 if lenA == 0 {
709 z.Set(b)
710 } else {
711 z.Set(a)
712 }
713 z.neg = false
714 if x != nil {
715 if lenA == 0 {
716 x.SetUint64(0)
717 } else {
718 x.SetUint64(1)
719 x.neg = negA
720 }
721 }
722 if y != nil {
723 if lenB == 0 {
724 y.SetUint64(0)
725 } else {
726 y.SetUint64(1)
727 y.neg = negB
728 }
729 }
730 return z
731 }
732
733 return z.lehmerGCD(x, y, a, b)
734 }
735
736
737
738
739
740
741
742
743
744
745
746
747
748 func lehmerSimulate(A, B *Int) (u0, u1, v0, v1 Word, even bool) {
749
750 var a1, a2, u2, v2 Word
751
752 m := len(B.abs)
753 n := len(A.abs)
754
755
756 h := nlz(A.abs[n-1])
757 a1 = A.abs[n-1]<<h | A.abs[n-2]>>(_W-h)
758
759 switch {
760 case n == m:
761 a2 = B.abs[n-1]<<h | B.abs[n-2]>>(_W-h)
762 case n == m+1:
763 a2 = B.abs[n-2] >> (_W - h)
764 default:
765 a2 = 0
766 }
767
768
769
770
771
772
773 even = false
774
775 u0, u1, u2 = 0, 1, 0
776 v0, v1, v2 = 0, 0, 1
777
778
779
780
781
782 for a2 >= v2 && a1-a2 >= v1+v2 {
783 q, r := a1/a2, a1%a2
784 a1, a2 = a2, r
785 u0, u1, u2 = u1, u2, u1+q*u2
786 v0, v1, v2 = v1, v2, v1+q*v2
787 even = !even
788 }
789 return
790 }
791
792
793
794
795
796
797
798
799
800
801 func lehmerUpdate(A, B, q, r *Int, u0, u1, v0, v1 Word, even bool) {
802 mulW(q, B, even, v0)
803 mulW(r, A, even, u1)
804 mulW(A, A, !even, u0)
805 mulW(B, B, !even, v1)
806 A.Add(A, q)
807 B.Add(B, r)
808 }
809
810
811
812 func mulW(z, x *Int, neg bool, w Word) {
813 z.abs = z.abs.mulAddWW(x.abs, w, 0)
814 z.neg = x.neg != neg
815 }
816
817
818
819
820 func euclidUpdate(A, B, Ua, Ub, q, r *Int, extended bool) (nA, nB, nr, nUa, nUb *Int) {
821 q.QuoRem(A, B, r)
822
823 if extended {
824
825 q.Mul(q, Ub)
826 Ua, Ub = Ub, Ua
827 Ub.Sub(Ub, q)
828 }
829
830 return B, r, A, Ua, Ub
831 }
832
833
834
835 type sixIntPool struct {
836 pool sync.Pool
837 }
838
839 func (t *sixIntPool) put(data *[6]Int) {
840 data[0].SetInt64(0)
841 data[1].SetInt64(0)
842 data[2].SetInt64(0)
843 data[3].SetInt64(0)
844 data[4].SetInt64(0)
845 data[5].SetInt64(0)
846 t.pool.Put(data)
847 }
848
849 func (t *sixIntPool) get() *[6]Int {
850 return t.pool.Get().(*[6]Int)
851 }
852
853 var sixIntP = sixIntPool{
854 sync.Pool{
855 New: func() any {
856 return &[6]Int{}
857 },
858 },
859 }
860
861
862
863
864
865
866
867
868
869
870
871 func (z *Int) lehmerGCD(x, y, a, b *Int) *Int {
872
873 data := sixIntP.get()
874 defer sixIntP.put(data)
875
876 var A, B, Ua, Ub *Int = &data[0], &data[1], &data[2], &data[3]
877
878 A.Abs(a)
879 B.Abs(b)
880
881 extended := x != nil || y != nil
882
883 if extended {
884
885 Ua.SetInt64(1)
886 }
887
888
889 q := &data[4]
890 r := &data[5]
891
892
893 if A.abs.cmp(B.abs) < 0 {
894 A, B = B, A
895 Ub, Ua = Ua, Ub
896 }
897
898
899 for len(B.abs) > 1 {
900
901 u0, u1, v0, v1, even := lehmerSimulate(A, B)
902
903
904 if v0 != 0 {
905
906
907
908 lehmerUpdate(A, B, q, r, u0, u1, v0, v1, even)
909
910 if extended {
911
912
913 lehmerUpdate(Ua, Ub, q, r, u0, u1, v0, v1, even)
914 }
915
916 } else {
917
918
919 A, B, r, Ua, Ub = euclidUpdate(A, B, Ua, Ub, q, r, extended)
920 }
921 }
922
923 if len(B.abs) > 0 {
924
925 if len(A.abs) > 1 {
926
927 A, B, r, Ua, Ub = euclidUpdate(A, B, Ua, Ub, q, r, extended)
928 }
929 if len(B.abs) > 0 {
930
931 aWord, bWord := A.abs[0], B.abs[0]
932 if extended {
933 var ua, ub, va, vb Word
934 ua, ub = 1, 0
935 va, vb = 0, 1
936 even := true
937 for bWord != 0 {
938 q, r := aWord/bWord, aWord%bWord
939 aWord, bWord = bWord, r
940 ua, ub = ub, ua+q*ub
941 va, vb = vb, va+q*vb
942 even = !even
943 }
944
945 mulW(Ua, Ua, !even, ua)
946 mulW(Ub, Ub, even, va)
947 Ua.Add(Ua, Ub)
948 } else {
949 for bWord != 0 {
950 aWord, bWord = bWord, aWord%bWord
951 }
952 }
953 A.abs[0] = aWord
954 }
955 }
956 negA := a.neg
957 if y != nil {
958
959 if y == b {
960 B.Set(b)
961 } else {
962 B = b
963 }
964
965 y.Mul(a, Ua)
966 if negA {
967 y.neg = !y.neg
968 }
969 y.Sub(A, y)
970 y.Div(y, B)
971 }
972
973 if x != nil {
974 x.Set(Ua)
975 if negA {
976 x.neg = !x.neg
977 }
978 }
979
980 z.Set(A)
981
982 return z
983 }
984
985
986
987
988
989 func (z *Int) Rand(rnd *rand.Rand, n *Int) *Int {
990
991 if n.neg || len(n.abs) == 0 {
992 z.neg = false
993 z.abs = nil
994 return z
995 }
996 z.neg = false
997 z.abs = z.abs.random(rnd, n.abs, n.abs.bitLen())
998 return z
999 }
1000
1001
1002
1003 type twoIntPool struct {
1004 pool sync.Pool
1005 }
1006
1007 func (t *twoIntPool) put(data *[2]Int) {
1008 data[0].SetInt64(0)
1009 data[1].SetInt64(0)
1010 t.pool.Put(data)
1011 }
1012
1013 func (t *twoIntPool) get() *[2]Int {
1014 return t.pool.Get().(*[2]Int)
1015 }
1016
1017 var twoIntP = twoIntPool{
1018 sync.Pool{
1019 New: func() any {
1020 return &[2]Int{}
1021 },
1022 },
1023 }
1024
1025
1026
1027
1028
1029 func (z *Int) ModInverse(g, n *Int) *Int {
1030
1031 if n.neg {
1032 var n2 Int
1033 n = n2.Neg(n)
1034 }
1035 if g.neg {
1036 var g2 Int
1037 g = g2.Mod(g, n)
1038 }
1039
1040
1041 data := twoIntP.get()
1042 defer twoIntP.put(data)
1043
1044 var d, x *Int = &data[0], &data[1]
1045 d.GCD(x, nil, g, n)
1046
1047
1048 if d.Cmp(intOne) != 0 {
1049 return nil
1050 }
1051
1052
1053
1054 if x.neg {
1055 z.Add(x, n)
1056 } else {
1057 z.Set(x)
1058 }
1059
1060 return z
1061 }
1062
1063 func (z nat) modInverse(g, n nat) nat {
1064
1065 return (&Int{abs: z}).ModInverse(&Int{abs: g}, &Int{abs: n}).abs
1066 }
1067
1068
1069
1070 func Jacobi(x, y *Int) int {
1071 if len(y.abs) == 0 || y.abs[0]&1 == 0 {
1072 panic(fmt.Sprintf("big: invalid 2nd argument to Int.Jacobi: need odd integer but got %s", y.String()))
1073 }
1074
1075
1076
1077
1078
1079 var a, b, c Int
1080 a.Set(x)
1081 b.Set(y)
1082 j := 1
1083
1084 if b.neg {
1085 if a.neg {
1086 j = -1
1087 }
1088 b.neg = false
1089 }
1090
1091 for {
1092 if b.Cmp(intOne) == 0 {
1093 return j
1094 }
1095 if len(a.abs) == 0 {
1096 return 0
1097 }
1098 a.Mod(&a, &b)
1099 if len(a.abs) == 0 {
1100 return 0
1101 }
1102
1103
1104
1105 s := a.abs.trailingZeroBits()
1106 if s&1 != 0 {
1107 bmod8 := b.abs[0] & 7
1108 if bmod8 == 3 || bmod8 == 5 {
1109 j = -j
1110 }
1111 }
1112 c.Rsh(&a, s)
1113
1114
1115 if b.abs[0]&3 == 3 && c.abs[0]&3 == 3 {
1116 j = -j
1117 }
1118 a.Set(&b)
1119 b.Set(&c)
1120 }
1121 }
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131 func (z *Int) modSqrt3Mod4Prime(x, p *Int) *Int {
1132 e := new(Int).Add(p, intOne)
1133 e.Rsh(e, 2)
1134 z.Exp(x, e, p)
1135 return z
1136 }
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146 func (z *Int) modSqrt5Mod8Prime(x, p *Int) *Int {
1147
1148
1149 e := new(Int).Rsh(p, 3)
1150 tx := new(Int).Lsh(x, 1)
1151 alpha := new(Int).Exp(tx, e, p)
1152 beta := new(Int).Mul(alpha, alpha)
1153 beta.Mod(beta, p)
1154 beta.Mul(beta, tx)
1155 beta.Mod(beta, p)
1156 beta.Sub(beta, intOne)
1157 beta.Mul(beta, x)
1158 beta.Mod(beta, p)
1159 beta.Mul(beta, alpha)
1160 z.Mod(beta, p)
1161 return z
1162 }
1163
1164
1165
1166 func (z *Int) modSqrtTonelliShanks(x, p *Int) *Int {
1167
1168 var s Int
1169 s.Sub(p, intOne)
1170 e := s.abs.trailingZeroBits()
1171 s.Rsh(&s, e)
1172
1173
1174 var n Int
1175 n.SetInt64(2)
1176 for Jacobi(&n, p) != -1 {
1177 n.Add(&n, intOne)
1178 }
1179
1180
1181
1182
1183
1184 var y, b, g, t Int
1185 y.Add(&s, intOne)
1186 y.Rsh(&y, 1)
1187 y.Exp(x, &y, p)
1188 b.Exp(x, &s, p)
1189 g.Exp(&n, &s, p)
1190 r := e
1191 for {
1192
1193 var m uint
1194 t.Set(&b)
1195 for t.Cmp(intOne) != 0 {
1196 t.Mul(&t, &t).Mod(&t, p)
1197 m++
1198 }
1199
1200 if m == 0 {
1201 return z.Set(&y)
1202 }
1203
1204 t.SetInt64(0).SetBit(&t, int(r-m-1), 1).Exp(&g, &t, p)
1205
1206 g.Mul(&t, &t).Mod(&g, p)
1207 y.Mul(&y, &t).Mod(&y, p)
1208 b.Mul(&b, &g).Mod(&b, p)
1209 r = m
1210 }
1211 }
1212
1213
1214
1215
1216
1217 func (z *Int) ModSqrt(x, p *Int) *Int {
1218 switch Jacobi(x, p) {
1219 case -1:
1220 return nil
1221 case 0:
1222 return z.SetInt64(0)
1223 case 1:
1224 break
1225 }
1226 if x.neg || x.Cmp(p) >= 0 {
1227 x = new(Int).Mod(x, p)
1228 }
1229
1230 switch {
1231 case p.abs[0]%4 == 3:
1232
1233 return z.modSqrt3Mod4Prime(x, p)
1234 case p.abs[0]%8 == 5:
1235
1236 return z.modSqrt5Mod8Prime(x, p)
1237 default:
1238
1239 return z.modSqrtTonelliShanks(x, p)
1240 }
1241 }
1242
1243
1244 func (z *Int) Lsh(x *Int, n uint) *Int {
1245 z.abs = z.abs.lsh(x.abs, n)
1246 z.neg = x.neg
1247 return z
1248 }
1249
1250
1251 func (z *Int) Rsh(x *Int, n uint) *Int {
1252 if x.neg {
1253
1254 t := z.abs.sub(x.abs, natOne)
1255 t = t.rsh(t, n)
1256 z.abs = t.add(t, natOne)
1257 z.neg = true
1258 return z
1259 }
1260
1261 z.abs = z.abs.rsh(x.abs, n)
1262 z.neg = false
1263 return z
1264 }
1265
1266
1267
1268 func (x *Int) Bit(i int) uint {
1269 if i == 0 {
1270
1271 if len(x.abs) > 0 {
1272 return uint(x.abs[0] & 1)
1273 }
1274 return 0
1275 }
1276 if i < 0 {
1277 panic("negative bit index")
1278 }
1279 if x.neg {
1280 t := nat(nil).sub(x.abs, natOne)
1281 return t.bit(uint(i)) ^ 1
1282 }
1283
1284 return x.abs.bit(uint(i))
1285 }
1286
1287
1288
1289
1290
1291
1292 func (z *Int) SetBit(x *Int, i int, b uint) *Int {
1293 if i < 0 {
1294 panic("negative bit index")
1295 }
1296 if x.neg {
1297 t := z.abs.sub(x.abs, natOne)
1298 t = t.setBit(t, uint(i), b^1)
1299 z.abs = t.add(t, natOne)
1300 z.neg = len(z.abs) > 0
1301 return z
1302 }
1303 z.abs = z.abs.setBit(x.abs, uint(i), b)
1304 z.neg = false
1305 return z
1306 }
1307
1308
1309 func (z *Int) And(x, y *Int) *Int {
1310 if x.neg == y.neg {
1311 if x.neg {
1312
1313 x1 := nat(nil).sub(x.abs, natOne)
1314 y1 := nat(nil).sub(y.abs, natOne)
1315 z.abs = z.abs.add(z.abs.or(x1, y1), natOne)
1316 z.neg = true
1317 return z
1318 }
1319
1320
1321 z.abs = z.abs.and(x.abs, y.abs)
1322 z.neg = false
1323 return z
1324 }
1325
1326
1327 if x.neg {
1328 x, y = y, x
1329 }
1330
1331
1332 y1 := nat(nil).sub(y.abs, natOne)
1333 z.abs = z.abs.andNot(x.abs, y1)
1334 z.neg = false
1335 return z
1336 }
1337
1338
1339 func (z *Int) AndNot(x, y *Int) *Int {
1340 if x.neg == y.neg {
1341 if x.neg {
1342
1343 x1 := nat(nil).sub(x.abs, natOne)
1344 y1 := nat(nil).sub(y.abs, natOne)
1345 z.abs = z.abs.andNot(y1, x1)
1346 z.neg = false
1347 return z
1348 }
1349
1350
1351 z.abs = z.abs.andNot(x.abs, y.abs)
1352 z.neg = false
1353 return z
1354 }
1355
1356 if x.neg {
1357
1358 x1 := nat(nil).sub(x.abs, natOne)
1359 z.abs = z.abs.add(z.abs.or(x1, y.abs), natOne)
1360 z.neg = true
1361 return z
1362 }
1363
1364
1365 y1 := nat(nil).sub(y.abs, natOne)
1366 z.abs = z.abs.and(x.abs, y1)
1367 z.neg = false
1368 return z
1369 }
1370
1371
1372 func (z *Int) Or(x, y *Int) *Int {
1373 if x.neg == y.neg {
1374 if x.neg {
1375
1376 x1 := nat(nil).sub(x.abs, natOne)
1377 y1 := nat(nil).sub(y.abs, natOne)
1378 z.abs = z.abs.add(z.abs.and(x1, y1), natOne)
1379 z.neg = true
1380 return z
1381 }
1382
1383
1384 z.abs = z.abs.or(x.abs, y.abs)
1385 z.neg = false
1386 return z
1387 }
1388
1389
1390 if x.neg {
1391 x, y = y, x
1392 }
1393
1394
1395 y1 := nat(nil).sub(y.abs, natOne)
1396 z.abs = z.abs.add(z.abs.andNot(y1, x.abs), natOne)
1397 z.neg = true
1398 return z
1399 }
1400
1401
1402 func (z *Int) Xor(x, y *Int) *Int {
1403 if x.neg == y.neg {
1404 if x.neg {
1405
1406 x1 := nat(nil).sub(x.abs, natOne)
1407 y1 := nat(nil).sub(y.abs, natOne)
1408 z.abs = z.abs.xor(x1, y1)
1409 z.neg = false
1410 return z
1411 }
1412
1413
1414 z.abs = z.abs.xor(x.abs, y.abs)
1415 z.neg = false
1416 return z
1417 }
1418
1419
1420 if x.neg {
1421 x, y = y, x
1422 }
1423
1424
1425 y1 := nat(nil).sub(y.abs, natOne)
1426 z.abs = z.abs.add(z.abs.xor(x.abs, y1), natOne)
1427 z.neg = true
1428 return z
1429 }
1430
1431
1432 func (z *Int) Not(x *Int) *Int {
1433 if x.neg {
1434
1435 z.abs = z.abs.sub(x.abs, natOne)
1436 z.neg = false
1437 return z
1438 }
1439
1440
1441 z.abs = z.abs.add(x.abs, natOne)
1442 z.neg = true
1443 return z
1444 }
1445
1446
1447
1448 func (z *Int) Sqrt(x *Int) *Int {
1449 if x.neg {
1450 panic("square root of negative number")
1451 }
1452 z.neg = false
1453 z.abs = z.abs.sqrt(nil, x.abs)
1454 return z
1455 }
1456
View as plain text