Ver Fonte

improve `random`

blueloveTH há 5 dias atrás
pai
commit
7edf139919
3 ficheiros alterados com 151 adições e 2 exclusões
  1. 25 0
      docs/modules/random.md
  2. 59 2
      src/modules/random.c
  3. 67 0
      tests/705_random.py

+ 25 - 0
docs/modules/random.md

@@ -30,3 +30,28 @@ Shuffle a sequence inplace.
 ### `random.choices(population, weights=None, k=1)`
 
 Return a k sized list of elements chosen from the population with replacement.
+
+### `random.getstate()`
+
+Return the internal state of the generator as a `bytes` object.
+
+### `random.setstate(state)`
+
+Restore the internal state of the generator from a `bytes` object
+returned by a previous call to `getstate()`.
+
+### `random.Random(x=None)`
+
+Create a new generator. `x` may be an `int` seed, `None` (seeded lazily from the
+system clock on first use), or a state returned by `getstate()`.
+
+A `Random` object supports `pickle`, so its state can be saved and restored:
+
+```python
+import pickle, random
+
+r = random.Random(7)
+data = pickle.dumps(r)
+r2 = pickle.loads(data)
+assert r.random() == r2.random()
+```

+ 59 - 2
src/modules/random.c

@@ -131,6 +131,25 @@ int64_t mt19937__randint(mt19937* self, int64_t a, int64_t b) {
     }
 }
 
+/* dumps the internal state as `bytes` */
+static void mt19937__getstate(mt19937* self, py_OutRef out) {
+    unsigned char* data = py_newbytes(out, sizeof(mt19937));
+    memcpy(data, self, sizeof(mt19937));
+}
+
+/* restores a state produced by `mt19937__getstate` */
+static bool mt19937__setstate(mt19937* self, py_Ref state) {
+    int size;
+    unsigned char* data = py_tobytes(state, &size);
+    if(size != sizeof(mt19937)) return ValueError("invalid state");
+    mt19937 tmp;
+    memcpy(&tmp, data, sizeof(mt19937));
+    /* `mti == N + 1` means mt[N] is not initialized; anything above `N` is out of range */
+    if(tmp.mti < 0 || tmp.mti > N + 1) return ValueError("invalid state");
+    *self = tmp;
+    return true;
+}
+
 static bool Random__new__(int argc, py_Ref argv) {
     mt19937* ud = py_newobject(py_retval(), py_totype(argv), 0, sizeof(mt19937));
     mt19937__ctor(ud);
@@ -142,13 +161,16 @@ static bool Random__init__(int argc, py_Ref argv) {
         // do nothing
     } else if(argc == 2) {
         mt19937* ud = py_touserdata(py_arg(0));
-        if(!py_isnone(&argv[1])) {
+        if(py_istype(py_arg(1), tp_bytes)) {
+            // a state returned by `getstate()`; this is how `__reduce__` rebuilds the object
+            if(!mt19937__setstate(ud, py_arg(1))) return false;
+        } else if(!py_isnone(&argv[1])) {
             PY_CHECK_ARG_TYPE(1, tp_int);
             py_i64 seed = py_toint(py_arg(1));
             mt19937__seed(ud, (uint32_t)seed);
         }
     } else {
-        return TypeError("Random(): expected 1 or 2 arguments, got %d");
+        return TypeError("Random(): expected 1 or 2 arguments, got %d", argc);
     }
     py_newnone(py_retval());
     return true;
@@ -169,6 +191,33 @@ static bool Random_seed(int argc, py_Ref argv) {
     return true;
 }
 
+static bool Random_getstate(int argc, py_Ref argv) {
+    PY_CHECK_ARGC(1);
+    mt19937* ud = py_touserdata(py_arg(0));
+    mt19937__getstate(ud, py_retval());
+    return true;
+}
+
+static bool Random_setstate(int argc, py_Ref argv) {
+    PY_CHECK_ARGC(2);
+    PY_CHECK_ARG_TYPE(1, tp_bytes);
+    mt19937* ud = py_touserdata(py_arg(0));
+    if(!mt19937__setstate(ud, py_arg(1))) return false;
+    py_newnone(py_retval());
+    return true;
+}
+
+static bool Random__reduce__(int argc, py_Ref argv) {
+    PY_CHECK_ARGC(1);
+    mt19937* ud = py_touserdata(py_arg(0));
+    // `(cls, (state,))`, i.e. `cls(state)` restores the generator
+    py_TValue* p = py_newtuple(py_retval(), 2);
+    py_assign(&p[0], py_tpobject(py_typeof(py_arg(0))));
+    py_TValue* args = py_newtuple(&p[1], 1);
+    mt19937__getstate(ud, &args[0]);
+    return true;
+}
+
 static bool Random_random(int argc, py_Ref argv) {
     PY_CHECK_ARGC(1);
     mt19937* ud = py_touserdata(py_arg(0));
@@ -298,9 +347,15 @@ void pk__add_module_random() {
     py_Ref mod = py_newmodule("random");
     py_Type type = py_newtype("Random", tp_object, mod, NULL);
 
+    // must be 2500 bytes so memcpy() works in `mt19937__getstate()` and `mt19937__setstate()`
+    _Static_assert(sizeof(mt19937) == 2500, "sizeof(mt19937) != 2500");
+
     py_bindmagic(type, __new__, Random__new__);
     py_bindmagic(type, __init__, Random__init__);
+    py_bindmagic(type, __reduce__, Random__reduce__);
     py_bindmethod(type, "seed", Random_seed);
+    py_bindmethod(type, "getstate", Random_getstate);
+    py_bindmethod(type, "setstate", Random_setstate);
     py_bindmethod(type, "random", Random_random);
     py_bindmethod(type, "uniform", Random_uniform);
     py_bindmethod(type, "randint", Random_randint);
@@ -317,6 +372,8 @@ void pk__add_module_random() {
     py_setdict(mod, py_name(name), py_retval());
 
     ADD_INST_BOUNDMETHOD("seed");
+    ADD_INST_BOUNDMETHOD("getstate");
+    ADD_INST_BOUNDMETHOD("setstate");
     ADD_INST_BOUNDMETHOD("random");
     ADD_INST_BOUNDMETHOD("uniform");
     ADD_INST_BOUNDMETHOD("randint");

+ 67 - 0
tests/705_random.py

@@ -73,3 +73,70 @@ assert c == randint(50, 100)
 
 import random
 assert random.Random(7).randint(1, 100) == a
+
+# test getstate/setstate
+r = random.Random(7)
+for _ in range(5):
+    r.random()
+
+state = r.getstate()
+assert isinstance(state, bytes)
+a = [r.randint(0, 1000) for _ in range(10)]
+r.setstate(state)
+assert a == [r.randint(0, 1000) for _ in range(10)]
+
+# a state can be moved between generators
+other = random.Random(123)
+other.setstate(r.getstate())
+assert other.random() == r.random()
+
+# `Random(state)` is equivalent to `setstate`
+assert random.Random(other.getstate()).random() == other.random()
+
+for bad in [b'', b'123', state[:-1]]:
+    try:
+        random.Random().setstate(bad)
+        exit(1)
+    except ValueError:
+        pass
+
+try:
+    random.Random().setstate(7)
+    exit(1)
+except TypeError:
+    pass
+
+# `mti` must stay within [0, 624+1]
+tmp = list(state)
+for mti in ([0xFF, 0xFF, 0xFF, 0xFF], [0x72, 0x02, 0, 0]):
+    try:
+        random.Random().setstate(bytes(tmp[:-4] + mti))
+        exit(1)
+    except ValueError:
+        pass
+
+# module-level generator exposes the same api
+random.seed(456)
+state = random.getstate()
+a = [random.random() for _ in range(5)]
+random.setstate(state)
+assert a == [random.random() for _ in range(5)]
+
+# test pickle
+import pickle
+
+r = random.Random(7)
+for _ in range(5):
+    r.random()
+
+r2 = pickle.loads(pickle.dumps(r))
+assert isinstance(r2, random.Random) and r2 is not r
+assert [r.random() for _ in range(10)] == [r2.random() for _ in range(10)]
+
+# an unseeded generator round-trips as unseeded
+fresh = random.Random()
+assert pickle.loads(pickle.dumps(fresh)).getstate() == fresh.getstate()
+
+# shared references are preserved
+res = pickle.loads(pickle.dumps([r, r]))
+assert res[0] is res[1]