[FIX] http: ensure values of session are serializable when stored

closes odoo/odoo#86015
Signed-off-by: Julien Castiaux <juc@odoo.com>
This commit is contained in:
Denis Ledoux
2023-02-07 09:12:32 +01:00
parent 40be0ebf9b
commit 536e670f8c
4 changed files with 113 additions and 20 deletions
@@ -1,6 +1,8 @@
# Part of Odoo. See LICENSE file for full copyright and licensing details.
import datetime
import json
import pytz
from urllib.parse import urlparse
from unittest.mock import patch
@@ -131,3 +133,76 @@ class TestHttpSession(TestHttpBase):
lang_fr.active = False
res = self.url_open('/test_http/echo-http-context-lang')
self.assertEqual(res.text, 'en_US')
def test_session7_serializable(self):
"""Tests setting a non-serializable value to the session is prevented
The test ensures the warning/exception is raised at the moment the attribute is set,
and not simply when the session is being saved in the session store.
"""
session = self.authenticate(None, None)
self.assertFalse(session.foo)
# Values allowed
for value in [
123,
12.3,
'foo',
(1, 2, 3, 4),
[1, 2, 3, 4],
set(),
{'1234'},
datetime.datetime.now(),
datetime.date.today(),
datetime.time(1, 33, 7),
pytz.timezone('UTC'),
pytz.timezone('Europe/Brussels'),
]:
session.foo = value
self.assertEqual(session.foo, value)
session.pop('foo')
self.assertFalse(session.foo)
session['foo'] = value
self.assertEqual(session.foo, value)
session.pop('foo')
# Values forbidden by odoo, raising a warning
for value in [
str,
int,
float,
bool,
range,
"foo".startswith,
datetime.datetime.strftime,
]:
with self.assertLogs(level="WARNING"):
session['foo'] = value
self.assertFalse(session.foo)
with self.assertLogs(level="WARNING"):
session.foo = value
self.assertFalse(session.foo)
with self.assertLogs(level="WARNING"):
# testing you cannot set a non-serializable value at the creation of the session
# e.g. in the __init__ of the session class
self.assertFalse(odoo.http.root.session_store.session_class({'foo': value}, 1234).foo)
with self.assertRaises(TypeError):
dict.update(session, foo=value)
self.assertFalse(session.foo)
# Values forbidden by pickle, raising an exception
for value in [
lambda: 'bar',
]:
with self.assertRaises(AttributeError):
session['foo'] = value
self.assertFalse(session.foo)
with self.assertRaises(AttributeError):
session.foo = value
self.assertFalse(session.foo)
with self.assertRaises(AttributeError):
# testing you cannot set a non-serializable value at the creation of the session
# e.g. in the __init__ of the session class
self.assertFalse(odoo.http.root.session_store.session_class({'foo': value}, 1234).foo)
with self.assertRaises(TypeError):
dict.update(session, foo=value)
self.assertFalse(session.foo)
+12 -9
View File
@@ -172,6 +172,7 @@ from .tools import (config, consteq, date_utils, file_path, parse_version,
profiler, submap, unique, ustr,)
from .tools.func import filter_kwargs, lazy_property
from .tools.mimetypes import guess_mimetype
from .tools.misc import pickle
from .tools._vendor import sessions
from .tools._vendor.useragents import UserAgent
@@ -893,12 +894,13 @@ class FilesystemSessionStore(sessions.FilesystemSessionStore):
class Session(collections.abc.MutableMapping):
""" Structure containing data persisted across requests. """
__slots__ = ('can_save', 'data', 'is_dirty', 'is_explicit', 'is_new',
__slots__ = ('can_save', '_Session__data', 'is_dirty', 'is_explicit', 'is_new',
'should_rotate', 'sid')
def __init__(self, data, sid, new=False):
self.can_save = True
self.data = data
self.__data = {}
self.update(data)
self.is_dirty = False
self.is_explicit = False
self.is_new = new
@@ -912,22 +914,23 @@ class Session(collections.abc.MutableMapping):
if item == 'geoip':
warnings.warn('request.session.geoip have been moved to request.geoip', DeprecationWarning)
return request.geoip if request else {}
return self.data[item]
return self.__data[item]
def __setitem__(self, item, value):
if item not in self.data or self.data[item] != value:
value = pickle.loads(pickle.dumps(value))
if item not in self.__data or self.__data[item] != value:
self.is_dirty = True
self.data[item] = value
self.__data[item] = value
def __delitem__(self, item):
del self.data[item]
del self.__data[item]
self.is_dirty = True
def __len__(self):
return len(self.data)
return len(self.__data)
def __iter__(self):
return iter(self.data)
return iter(self.__data)
def __getattr__(self, attr):
return self.get(attr, None)
@@ -939,7 +942,7 @@ class Session(collections.abc.MutableMapping):
self[key] = val
def clear(self):
self.data.clear()
self.__data.clear()
self.is_dirty = True
#
+3 -5
View File
@@ -21,9 +21,7 @@ import re
import tempfile
from hashlib import sha1
from os import path, replace as rename
from pickle import dump
from pickle import HIGHEST_PROTOCOL
from pickle import load
from odoo.tools.misc import pickle
from time import time
from werkzeug.datastructures import CallbackDict
@@ -195,7 +193,7 @@ class FilesystemSessionStore(SessionStore):
fd, tmp = tempfile.mkstemp(suffix=_fs_transaction_suffix, dir=self.path)
f = os.fdopen(fd, "wb")
try:
dump(dict(session), f, HIGHEST_PROTOCOL)
pickle.dump(dict(session), f, pickle.HIGHEST_PROTOCOL)
finally:
f.close()
try:
@@ -224,7 +222,7 @@ class FilesystemSessionStore(SessionStore):
else:
try:
try:
data = load(f)
data = pickle.load(f, errors={})
except Exception:
_logger.debug('Could not load session data. Use empty session.', exc_info=True)
data = {}
+23 -6
View File
@@ -1581,15 +1581,31 @@ def format_duration(value):
consteq = hmac_lib.compare_digest
_PICKLE_SAFE_NAMES = {
'builtins': [
'set', # Required to support `set()` for Python < 3.8
],
'datetime': [
'datetime',
'date',
'time',
],
'pytz': [
'_p',
'_UTC',
],
}
# https://docs.python.org/3/library/pickle.html#restricting-globals
# forbid globals entirely: str/unicode, int/long, float, bool, tuple, list, dict, None
class Unpickler(pickle_.Unpickler, object):
find_global = None # Python 2
find_class = None # Python 3
def find_class(self, module_name, name):
safe_names = _PICKLE_SAFE_NAMES.get(module_name, [])
if name in safe_names:
return super().find_class(module_name, name)
raise AttributeError("global '%s.%s' is forbidden" % (module_name, name))
def _pickle_load(stream, encoding='ASCII', errors=False):
if sys.version_info[0] == 3:
unpickler = Unpickler(stream, encoding=encoding)
else:
unpickler = Unpickler(stream)
unpickler = Unpickler(stream, encoding=encoding)
try:
return unpickler.load()
except Exception:
@@ -1601,6 +1617,7 @@ pickle.load = _pickle_load
pickle.loads = lambda text, encoding='ASCII': _pickle_load(io.BytesIO(text), encoding=encoding)
pickle.dump = pickle_.dump
pickle.dumps = pickle_.dumps
pickle.HIGHEST_PROTOCOL = pickle_.HIGHEST_PROTOCOL
class DotDict(dict):